// SIMD vector optimization rules

///|
/// Helper to create splat64 pattern: replicate 64-bit value 2 times
fn splat64(n : UInt64) -> Bytes {
  let val = n.reinterpret_as_int64()
  Bytes::makei(16, fn(i) {
    let lane = i / 8
    let byte_idx = i % 8
    guard lane < 2 else { b'\x00' }
    ((val >> (byte_idx * 8)) & 0xFFL).to_byte()
  })
}

///|
/// splat(iconst) -> vconst
/// Converts splat of constant to vector constant
fn rule_splat_const() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Splat &&
          node.children.length() == 1 &&
          eg.find_const(node.children[0]) is Some(c) {
          // Default to 64-bit splat (most common case)
          // In a real implementation, we'd need type information
          // to determine the lane width
          let pattern = splat64(c.reinterpret_as_uint64())
          let vconst_node = eg.add_vconst(pattern)
          changed = eg.merge_changed(class_id, vconst_node) || changed
        }
      }
      changed
    },
  }
}

// ============================================================================
// Lift splat outside of binary operations
// op(splat(x), splat(y)) -> splat(op(x, y))
// ============================================================================

///|
/// Helper to check if a node is a Splat and get its child
fn find_splat_child(eg : EGraph, class_id : EClassId) -> EClassId? {
  for node in eg.get_nodes(class_id) {
    if node.op is Splat && node.children.length() == 1 {
      return Some(node.children[0])
    }
  }
  None
}

///|
/// band(splat(x), splat(y)) -> splat(band(x, y))
fn rule_band_splat_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is And &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) &&
          find_splat_child(eg, node.children[1]) is Some(y) {
          // band(splat(x), splat(y)) -> splat(band(x, y))
          let inner_and = eg.add_and(x, y)
          let new_node = eg.add({ op: Splat, children: [inner_and] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// bor(splat(x), splat(y)) -> splat(bor(x, y))
fn rule_bor_splat_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Or &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) &&
          find_splat_child(eg, node.children[1]) is Some(y) {
          let inner_or = eg.add_or(x, y)
          let new_node = eg.add({ op: Splat, children: [inner_or] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// bxor(splat(x), splat(y)) -> splat(bxor(x, y))
fn rule_bxor_splat_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Xor &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) &&
          find_splat_child(eg, node.children[1]) is Some(y) {
          let inner_xor = eg.add_xor(x, y)
          let new_node = eg.add({ op: Splat, children: [inner_xor] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// bnot(splat(x)) -> splat(bnot(x))
fn rule_bnot_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Bnot &&
          node.children.length() == 1 &&
          find_splat_child(eg, node.children[0]) is Some(x) {
          let inner_bnot = eg.add({ op: Bnot, children: [x] })
          let new_node = eg.add({ op: Splat, children: [inner_bnot] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// iadd(splat(x), splat(y)) -> splat(iadd(x, y))
fn rule_iadd_splat_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Add &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) &&
          find_splat_child(eg, node.children[1]) is Some(y) {
          let inner_add = eg.add_add(x, y)
          let new_node = eg.add({ op: Splat, children: [inner_add] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// isub(splat(x), splat(y)) -> splat(isub(x, y))
fn rule_isub_splat_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Sub &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) &&
          find_splat_child(eg, node.children[1]) is Some(y) {
          let inner_sub = eg.add_sub(x, y)
          let new_node = eg.add({ op: Splat, children: [inner_sub] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// imul(splat(x), splat(y)) -> splat(imul(x, y))
fn rule_imul_splat_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Mul &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) &&
          find_splat_child(eg, node.children[1]) is Some(y) {
          let inner_mul = eg.add_mul(x, y)
          let new_node = eg.add({ op: Splat, children: [inner_mul] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// ineg(splat(x)) -> splat(ineg(x))
fn rule_ineg_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Neg &&
          node.children.length() == 1 &&
          find_splat_child(eg, node.children[0]) is Some(x) {
          let inner_neg = eg.add({ op: Neg, children: [x] })
          let new_node = eg.add({ op: Splat, children: [inner_neg] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// iabs(splat(x)) -> splat(iabs(x))
fn rule_iabs_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Iabs &&
          node.children.length() == 1 &&
          find_splat_child(eg, node.children[0]) is Some(x) {
          let inner_iabs = eg.add({ op: Iabs, children: [x] })
          let new_node = eg.add({ op: Splat, children: [inner_iabs] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// popcnt(splat(x)) -> splat(popcnt(x))
fn rule_popcnt_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Popcnt &&
          node.children.length() == 1 &&
          find_splat_child(eg, node.children[0]) is Some(x) {
          let inner_popcnt = eg.add({ op: Popcnt, children: [x] })
          let new_node = eg.add({ op: Splat, children: [inner_popcnt] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// smin(splat(x), splat(y)) -> splat(smin(x, y))
fn rule_smin_splat_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Smin &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) &&
          find_splat_child(eg, node.children[1]) is Some(y) {
          let inner_smin = eg.add({ op: Smin, children: [x, y] })
          let new_node = eg.add({ op: Splat, children: [inner_smin] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// umin(splat(x), splat(y)) -> splat(umin(x, y))
fn rule_umin_splat_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Umin &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) &&
          find_splat_child(eg, node.children[1]) is Some(y) {
          let inner_umin = eg.add({ op: Umin, children: [x, y] })
          let new_node = eg.add({ op: Splat, children: [inner_umin] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// smax(splat(x), splat(y)) -> splat(smax(x, y))
fn rule_smax_splat_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Smax &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) &&
          find_splat_child(eg, node.children[1]) is Some(y) {
          let inner_smax = eg.add({ op: Smax, children: [x, y] })
          let new_node = eg.add({ op: Splat, children: [inner_smax] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// umax(splat(x), splat(y)) -> splat(umax(x, y))
fn rule_umax_splat_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Umax &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) &&
          find_splat_child(eg, node.children[1]) is Some(y) {
          let inner_umax = eg.add({ op: Umax, children: [x, y] })
          let new_node = eg.add({ op: Splat, children: [inner_umax] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

// ============================================================================
// Shift/rotate operations: only first operand is splatted
// op(splat(x), y) -> splat(op(x, y))
// ============================================================================

///|
/// rotl(splat(x), y) -> splat(rotl(x, y))
fn rule_rotl_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Rotl &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) {
          let y = node.children[1]
          let inner_rotl = eg.add({ op: Rotl, children: [x, y] })
          let new_node = eg.add({ op: Splat, children: [inner_rotl] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// rotr(splat(x), y) -> splat(rotr(x, y))
fn rule_rotr_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Rotr &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) {
          let y = node.children[1]
          let inner_rotr = eg.add({ op: Rotr, children: [x, y] })
          let new_node = eg.add({ op: Splat, children: [inner_rotr] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// ishl(splat(x), y) -> splat(ishl(x, y))
fn rule_ishl_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Shl &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) {
          let y = node.children[1]
          let inner_shl = eg.add_shl(x, y)
          let new_node = eg.add({ op: Splat, children: [inner_shl] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// ushr(splat(x), y) -> splat(ushr(x, y))
fn rule_ushr_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Ushr &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) {
          let y = node.children[1]
          let inner_ushr = eg.add({ op: Ushr, children: [x, y] })
          let new_node = eg.add({ op: Splat, children: [inner_ushr] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}

///|
/// sshr(splat(x), y) -> splat(sshr(x, y))
fn rule_sshr_splat() -> RewriteRule {
  {
    apply: fn(eg, class_id) {
      let mut changed = false
      for node in eg.get_nodes(class_id) {
        if node.op is Sshr &&
          node.children.length() == 2 &&
          find_splat_child(eg, node.children[0]) is Some(x) {
          let y = node.children[1]
          let inner_sshr = eg.add({ op: Sshr, children: [x, y] })
          let new_node = eg.add({ op: Splat, children: [inner_sshr] })
          changed = eg.merge_changed(class_id, new_node) || changed
        }
      }
      changed
    },
  }
}