from dataclasses import replace
from tinygrad.dtype import dtypes, to_dtype
from tinygrad.uop.ops import PatternMatcher, UPat, Ops, UOp, resolve, GroupOp, ParamArg
from tinygrad.uop.ops import graph_rewrite, rewrite_group, identity_element, resolve_returned_after
from tinygrad.uop.movement import mop_cleanup
from tinygrad.helpers import prod, getenv, all_int, DEBUG, SPLIT_REDUCEOP, OPENPILOT_HACKS, FLOAT16, argsort
from tinygrad.schedule.indexing import apply_movement_op
from tinygrad.schedule.allreduce import create_allreduce_function
from tinygrad.schedule.multi import multi_pm

def forward_call_outputs(sink:UOp) -> UOp:
  placed:dict[UOp, UOp] = {}
  items:list[UOp] = []
  for item in sink.src:
    st = item.src[1] if item.op is Ops.AFTER and len(item.src) == 2 and item.src[1].op is Ops.STORE else item
    if st.op is not Ops.STORE or (item is not st and item.src[0] is not st.src[0]):
      items.append(item)
      continue
    target, src = st.src
    while src.op is Ops.AFTER: src = src.src[0]
    base = src.storage_base
    if item is not st and base.op is not Ops.ALLOC:
      items.append(item)
      continue
    # Forward the allocation, not just one view of it, so saved values and other aliases follow the same placement.
    key = base if base.op is Ops.ALLOC else src
    if key not in placed and (src.op is Ops.STAGE or target.has_buffer_identity()) and \
       target.storage_base not in st.src[1].toposort(enter_calls=False):
      if base.op is Ops.ALLOC and src.has_buffer_identity() and base.max_numel() == target.storage_base.max_numel():
        placed[key] = target.storage_base
      elif src.op is Ops.STAGE: placed[key] = target.after(target.store(src.src[0]))
      elif src.op in {Ops.BUFFER, Ops.UNSHARD} and src.has_buffer_identity(): placed[key] = target
      if key in placed:
        if item is not st: placed[item] = st.src[1]
        items.append(st.src[1])
        continue
    items.append(target.after(st))
  return UOp.sink(*items).substitute(placed, walk=True)

def walk_mop(u:UOp):
  if u.op in GroupOp.Movement or u.op in {Ops.INDEX, Ops.UNSHARD, Ops.BITCAST}: return walk_mop(u.src[0])
  if u.op is Ops.AFTER and (b:=walk_mop(u.src[0])) is not u.src[0]: return b.after(*u.src[1:])
  return u

def found_after(ctx:dict[UOp, UOp], after:UOp, src:UOp):
  if (x:=src).op is Ops.CAST and x.dtype == dtypes.half and FLOAT16: x, after = x.src[0], after.cast(dtypes.float)
  while True:
    if x.op is Ops.PERMUTE: x, after = x.src[0], after.permute(argsort(x.marg))
    elif x.op is Ops.RESHAPE: x, after = x.src[0], after.reshape(x.src[0].shape)
    elif x.op is Ops.WHERE and x.src[2].base.is_invalid and x.src[1].op is Ops.PAD:
      x, after = x.src[1].src[0], after.shrink(tuple((o, s+o) for (o,_),s in zip(x.src[1].marg, x.src[1].src[0].shape)))
    else: break
  ctx[x] = after

# *** fold moved AFTERs (hack for openpilot) ***
pm_fold_moved_after = PatternMatcher([
  (UPat(Ops.AFTER, src=(UPat(), UPat(Ops.STORE, src=(UPat(), UPat((*GroupOp.Movement,Ops.CAST,Ops.WHERE), name="src")))), name="after"), found_after),
  # contiguous is also a materialization point (it bufferizes in the scheduler)
  (UPat(Ops.STAGE, src=(UPat((*GroupOp.Movement,Ops.CAST,Ops.WHERE), name="src"),), name="after"), found_after),
  # replace ALU sources with AFTER versions found above
  (UPat(GroupOp.ALU, name="alu"), lambda ctx,alu: alu.replace(src=new_src) if (new_src:=tuple(ctx.get(s, s) for s in alu.src)) != alu.src else None),
])

# movement op on INDEX as a PatternMatcher
def _mop_index(r:UOp, idx:UOp):
  idxs = idx.src[1:]
  if len(idxs) == len(r.shape):
    return r.src[0].index(*apply_movement_op(r.op, r.src[0].shape, r.marg, idxs), arg=idx.arg)
  if r.op is Ops.RESHAPE:
    src_prefix = len(r.src[0].shape) - len(r.shape[len(idxs):])
    if src_prefix >= 0 and r.src[0].shape[src_prefix:] == r.shape[len(idxs):]:
      if src_prefix == 0: return r.src[0]
      ret = r.src[0].index(*apply_movement_op(r.op, r.src[0].shape[:src_prefix], r.shape[:len(idxs)], idxs), arg=idx.arg)
      return ret if ret.shape == idx.shape else None

pm_mops = PatternMatcher([
  # handle movement ops on INDEX
  (UPat(GroupOp.Movement, name="r").f(Ops.INDEX, allow_any_len=True, name="idx"), _mop_index),
  # move movement ops and INDEX after AFTER
  (UPat(GroupOp.Movement|{Ops.INDEX}, name="r").after(name="a", allow_any_len=True),
   lambda r,a: UOp(r.op, src=(a.replace(src=(r.src[0],)+a.src[1:]),)+r.src[1:], arg=r.arg)),
  (UPat(GroupOp.Movement, name="r").end(name="a", allow_any_len=True), lambda r,a: a.replace(src=(r.src[0],)+a.src[1:])),
])

# *****************
# 0. do some cleanup rewrites, mostly copied from the old stuff

# stop at materialization boundaries, including COPYs/STAGEs already lowered to AFTER+STORE
def store_hazard_boundary(s:UOp):
  if s.op in {Ops.COPY, Ops.STAGE}: return False
  if s.op is Ops.AFTER: return not any(d.op is Ops.STORE and d.src[0].base is s.src[0].base for d in s.src[1:])
  return True

def fix_store_hazard(target:UOp, src:UOp):
  if (base:=target.base) not in src.toposort(enter_calls=False): return None
  # PERMUTE and FLIP reorder indices, SHRINK can have overlapping regions when dest is also shrunk
  unsafe = {Ops.PERMUTE, Ops.FLIP} | ({Ops.SHRINK} if target.op_in_backward_slice_with_self(Ops.SHRINK) else set())
  reaches_base: dict[UOp, bool] = {}
  for s in src.toposort(gate=store_hazard_boundary):
    reaches_base[s] = s is base or any(reaches_base.get(c) for c in s.src)
    if reaches_base[s] and s.op in unsafe and not (s is target and s.op is Ops.SHRINK): return target.store(src.contiguous())

def split_reduceop(reduce:UOp, x:UOp):
  if prod(reduce.shape) == 0: return None
  if not SPLIT_REDUCEOP or not all_int(x.shape) or (prod(x.shape)//prod(reduce.shape))<getenv("REDUCEOP_SPLIT_THRESHOLD", 32768): return None
  # if there are few globals, make some reduces into globals by splitting into two kernels
  # cap output buffer to 2**22: heuristic number of global outputs to achieve max occupancy with enough locals+upcasts for gemm
  #   ~2**10 should be enough if GROUP is used
  # 256 split maximum should be "negligible reduce" for low prod(reduce.shape), 8 split minimum.
  # split is moved to the end to provide maximum locality for the second phase reduce.

  # get expanded by rangeifying the UOp x
  indexed = x.index(*[UOp.range(s, i) if resolve(s>1) else 0 for i,s in enumerate(x.shape)])
  range_nums = [y.axis_id[0] for y in indexed.substitute({x.base:UOp(Ops.NOOP)}, extra_pm=pm_mops).ranges]
  is_expanded = [i not in range_nums for i in range(len(x.shape))]

  if not (split_candidates:=[(i,d) for i in range(reduce.arg[1])
                             for d in range(min(256,2**getenv("REDUCEOP_SPLIT_SIZE",22)//prod(reduce.shape)),8-1,-1)
                             if x.shape[i]%d==0 and not is_expanded[i]]): return None
  dim_to_split, divisor = split_candidates[0]
  splitted_shape = x.shape[:dim_to_split]+(divisor,)+(x.shape[dim_to_split]//divisor,)+x.shape[dim_to_split+1:]
  splitted = x.reshape(splitted_shape).permute(tuple([d for d in range(len(splitted_shape)) if d!=dim_to_split]+[dim_to_split]))
  if DEBUG >= 3: print(f"split {divisor}: {x.shape} -> {splitted.shape} -> {reduce.shape}")
  # reduce original axes, then split
  return splitted._rop(reduce.arg[0], tuple(range(reduce.arg[1]))).contiguous()._rop(reduce.arg[0], (len(reduce.shape),))

def resolve_function(c:UOp) -> UOp|None:
  if not c.is_inline_call: return None
  nodes = c.body.toposort(enter_calls=False)
  # Input and output PARAMs both bind to explicit arguments by slot; unused arguments are allowed.
  args = c.src[1:]

  # params have a flat storage size in the arg, the logical shape is a view (RESHAPE/SHRINK/UNSHARD) on top of it.
  # substitute args by their flat max-shaped storage view so the movement views on the params stay valid
  def flat_storage(a:UOp) -> tuple[int, UOp]:  # returns (size, view of a as flat max-shaped storage)
    shp = a.max_shard_shape if a.axis is not None and isinstance(a.device, tuple) else a.max_shape
    if a.op is Ops.SHRINK and a.src[0].shape == shp and all(s == 0 for s,_ in a.marg): a = a.src[0]
    return (n:=prod(shp)), a if a.shape == (n,) else a.pad_to(shp).reshape((n,))
  dict_map = {p:args[p.arg.slot] for p in nodes if p.op is Ops.PARAM and p.arg.slot >= 0}
  for p, a in dict_map.items():
    if p.shape:
      n, flat = flat_storage(a)
      if p.src[0].val != n: raise TypeError(f"arg {p.arg.slot} shape mismatch: expected size {p.src[0].val}, got {a.shape}")
      dict_map[p] = flat
    elif a.shape != ():
      raise TypeError(f"arg {p.arg.slot} shape mismatch: expected scalar, got {a.shape}")
    if p.dtype != a.dtype: raise TypeError(f"arg {p.arg.slot} dtype mismatch: expected {p.dtype}, got {a.dtype}")
  # Inlining removes the call scope, so its local allocations need fresh identities.
  dict_map.update({b:b.replace(arg=replace(b.arg, slot=next(UOp.unique_num))) for b in nodes if b.op is Ops.ALLOC})
  return c.body.substitute(dict_map, walk=True)

# shape-changing bitcast
def expand_bitcast(bc:UOp) -> UOp|None:
  x = bc.src[0]
  if (ns:=bc.dtype.itemsize) == (os:=x.dtype.itemsize) or x.on_disk(): return None
  new_uint, tmp = to_dtype(f"uint{8*ns}"), x.bitcast(to_dtype(f"uint{8*os}"))
  if ns > os:
    tmp = tmp.reshape(x.shape[:-1] + (x.shape[-1]//(rate := ns//os), rate))
    parts = [tmp.shrink((None,)*(len(tmp.shape)-1) + ((i, i+1),)).cast(new_uint)<<8*i*os for i in range(rate)]
    return parts[0].usum(*parts[1:]).squeeze(-1).bitcast(bc.dtype)
  parts = [tmp>>8*i*ns for i in range(os//ns)]
  return parts[0].stack(*parts[1:], dim=-1).flatten(-2).cast(new_uint).bitcast(bc.dtype)

def copy_to_anon_store(x:UOp, copy:UOp):
  # copies are always cross device: pad to the max shape so the copy reads a whole buffer (SDMA can't do offset copies)
  x = x.pad_to(x.max_shape)
  # the buffer takes the DEVICE range from the copy (no-op for single device copies)
  buf = UOp(Ops.ALLOC, src=(UOp.const(prod(x.max_shape)),)+copy.src[1:],
            arg=ParamArg(next(UOp.unique_num), copy.dtype, device=copy.device)).reshape(x.max_shape)
  return buf.after(buf.store(x)).shrink_to(copy.shape)

def stage_to_anon_store(x:UOp, stg:UOp):
  # the buffer created here is inside the call and is not persisted, like the buffers created for copies
  buf = UOp(Ops.ALLOC, src=(UOp.const(prod(x.max_shape)),)+UOp.device_range_src(x.device),
            arg=ParamArg(next(UOp.unique_num), stg.dtype, device=x.device)).reshape(x.max_shape)
  view = buf.shrink_to(stg.shape)
  return view.after(view.store(x))

def materialize_cross_device_src(dest:UOp, src:UOp):
  # cross-device copies must read a whole buffer (SDMA can't do offset copies)
  if src.device is None or dest.device == src.device or src.has_buffer_identity(after_ok=True): return None
  return dest.store(src.contiguous())

pm_inline_calls = PatternMatcher([
  (UPat(Ops.CALL, name="c"), resolve_function),
  (UPat(Ops.AFTER, src=(UPat(name="r"), UPat(Ops.SINK, name="t")), allow_any_len=True), resolve_returned_after),
])

pm_disk_copy = PatternMatcher([
  # remove contiguous on movement ops before a copy on disk
  (UPat(GroupOp.Movement, name="x").f(Ops.STAGE).f(Ops.COPY, name="copy"), lambda x,copy:
   copy.replace(src=(x,)) if x.on_disk() else None),
  # push all movement ops to the destination: views exposed here are no longer normalized into input PARAMs,
  # so leaving SHRINK/RESHAPE behind can cause materialize_cross_device_src to allocate a temporary on disk
  (UPat(GroupOp.Movement, name="x").f(Ops.COPY, name="copy"), lambda x,copy:
   x.replace(src=(copy.replace(src=(x.src[0],)),)+x.src[1:]) if x.on_disk() else None),
])

earliest_rewrites = mop_cleanup+PatternMatcher([
  # resolve allreduce (must be bottom up)
  (UPat(Ops.ALLREDUCE, src=(UPat.var("buf"),), name="red"), create_allreduce_function),

  # split_reduceop
  (UPat(Ops.REDUCE, name="reduce", src=(UPat.var("x"),)), split_reduceop),

  # remove DETACH/CONTIGUOUS_BACKWARD (TODO: this is copied in allocations)
  (UPat((Ops.DETACH, Ops.CONTIGUOUS_BACKWARD), name="x"), lambda x: x.src[0]),

  # SINK only ever references the base
  (UPat(Ops.SINK, name="x"), lambda x: x.replace(src=tuple(y.unsharded_base for y in x.src))),

  # ** copy rules **

  # a copy to the same device as the source is not allowed: it is a no-op, STAGE materializes on the same device
  (UPat(Ops.COPY, src=(UPat.var("x"),), allow_any_len=True, name="copy"), lambda x,copy: x if x.device == copy.device else None),

  # a COPY in src[1] of a plain STORE can just be removed: a STORE to a buffer on a different device is a COPY
  (UPat(Ops.STORE, src=(UPat.var("dst"), UPat(Ops.COPY, src=(UPat.var("x"),), allow_any_len=True, name="cpy"))),
   lambda dst,x,cpy: dst.store(x) if dst.device == cpy.device and dst.has_buffer_identity(after_ok=True) else None),

  # a bare COPY is an anonymous store: realize it as a STORE into a fresh call-local buffer on the copy device
  (UPat(Ops.COPY, src=(UPat.var("x"),), allow_any_len=True, name="copy"), copy_to_anon_store),

  # ** stage rules **

  # a STAGE of an already materialized value (or of a COPY, which materializes itself) is a no-op
  (UPat(Ops.STAGE, src=(UPat.var("x"),)),
   lambda x: x if x.has_buffer_identity(after_ok=True) or x.op is Ops.COPY else None),

  # a bare STAGE is an anonymous same-device materialization: realize it as a STORE into a fresh call-local buffer
  (UPat(Ops.STAGE, src=(UPat.var("x"),), name="stg"), stage_to_anon_store),

  # reshaping on STORE can be a NOOP
  (UPat(Ops.STORE, src=(UPat(Ops.RESHAPE, src=(UPat.var("dst",),), allow_any_len=True),
                        UPat(Ops.RESHAPE, src=(UPat.var("src",),), allow_any_len=True))),
   lambda dst,src: dst.store(src) if dst.shape == src.shape else None),

  # ** store rules **

  # materialize the src of a cross device STORE on its own device first: the STORE itself is the copy
  (UPat(Ops.STORE, src=(UPat(name="dest"), UPat(name="src"))), materialize_cross_device_src),

  # fix store hazard (dest is in used in src) by adding contiguous: TestAssign.test_post_flipped_assignment
  (UPat(Ops.STORE, src=(UPat(name="target"), UPat(name="src"))), fix_store_hazard),

  # remove two STOREs that store the same thing to the same place: TestSchedule.test_dedup_Assign
  (UPat.var("buf").after(UPat.var("buf").store(UPat.var("src")), name="a1").after(UPat.var("a1").store(UPat.var("src"))), lambda buf,src,a1:a1),

  # store a buffer's own current contents back into itself: TestAssign.test_assign_from_alias
  (UPat.var("buf").after(UPat.var("buf").store(UPat.var("buf").after(UPat.var("buf").store(UPat()), name="a1"))), lambda buf,a1:a1),

  # move bitcast from store dest to source: TestAssign.test_assign_bitcast
  (UPat(Ops.STORE, src=(UPat(Ops.BITCAST, src=(UPat(name="target"),)), UPat(name="src"))),
   lambda target, src: target.store(src.bitcast(target.dtype))),

  (UPat(Ops.BITCAST, name="bc"), expand_bitcast),

  # ** size 0 **

  # reduce of size 0 is the identity element
  (UPat(Ops.REDUCE, name="reduce", src=(UPat.var("x"),)),
   lambda reduce,x: reduce.const_like(identity_element(reduce.arg[0], reduce.dtype)) if 0 in x.shape and 0 not in reduce.shape else None),
  # handle size 0
  (UPat(GroupOp.All-{Ops.SINK}, name="x"), lambda x: x.const_like(0).rtag(x.tag) if x._shape is not None and 0 in x.shape else None),

  # remove movement ops from SINK/AFTER. TODO: should be generic
  (UPat(Ops.SINK, name="s"), lambda s: s.replace(src=tuple(walk_mop(u) for u in s.src if u.op is not Ops.NOOP))),
  (UPat(Ops.AFTER, name="s"), lambda s: s.replace(src=(s.src[0],)+tuple(walk_mop(u) for u in s.src[1:] if u.op is not Ops.NOOP))),
])

@rewrite_group(new_ctx=False)
def prepare_rangeify(sink:UOp) -> UOp:
  # prepare for rangeify
  tsink = graph_rewrite(forward_call_outputs(sink), multi_pm, name="multi_pm")
  tsink = graph_rewrite(tsink, pm_mops+pm_inline_calls+pm_disk_copy, name="inline calls")
  if OPENPILOT_HACKS: tsink = graph_rewrite(tsink, pm_fold_moved_after, ctx={}, name="fold moved afters")
  tsink = graph_rewrite(tsink, pm_mops+earliest_rewrites, bottom_up=True, name="earliest rewrites")
  return tsink
