import {
  aGtB,
  checkArrayOfNumbers,
  checkNumber,
  funArrAbs,
  funArrAdd,
  funArrDiv,
  funArrSub,
  funMA,
  funRound,
  funMS,
} from "../common";

/**
 * DMI - Directional Movement Index
 * @param {number[]} CLOSE - close prices
 * @param {number[]} HIGH - high prices
 * @param {number[]} LOW - low prices
 * @param {number} N - number of periods,default 14
 * @param {number} M - number of periods for moving average,default 6
 * @returns {Object} result
 */

export function DMI(
  CLOSE: number[],
  HIGH: number[],
  LOW: number[],
  N: number = 14,
  M: number = 6
) {
  // - check input data
  checkArrayOfNumbers(CLOSE, "CLOSE");
  checkArrayOfNumbers(HIGH, "HIGH");
  checkArrayOfNumbers(LOW, "LOW");

  checkNumber(N, 2, CLOSE, "N", "CLOSE");
  checkNumber(M, 2, N, "M", "N");

  // - calculate DMI
  const tmpM = CLOSE.slice(1).reduce(
    (acc: number[], cur, index) => {
      acc.push(
        Math.max(
          Math.max(
            HIGH[index + 1] - LOW[index + 1],
            Math.abs(HIGH[index + 1] - CLOSE[index])
          ),
          Math.abs(CLOSE[index]) - LOW[index + 1]
        )
      );
      return acc;
    },
    [0]
  );

  const MTR = funMS(tmpM, N);
  const [HD, LD] = CLOSE.slice(1).reduce(
    (acc: [number[], number[]], cur, index) => {
      acc[0].push(HIGH[index + 1] - HIGH[index]);
      acc[1].push(LOW[index] - LOW[index + 1]);
      return acc;
    },
    [[0], [0]]
  );

  const HDM = HD.map((v, i) => (v > 0 && aGtB(v, LD[i]) ? v : 0)); //浮点数比较大小，使用aGtB函数
  const LDM = LD.map((v, i) => (v > 0 && aGtB(v, HD[i]) ? v : 0));

  const DMP = funMS(HDM, N);
  const DMM = funMS(LDM, N);
  //以上数据正确

  const PDI: number[] = [];
  const MDI: number[] = [];

  for (let i = 0; i < DMP.length; i++) {
    PDI.push(((DMP[i] ?? 0) * 100) / (MTR[i] ?? 1));
    MDI.push(((DMM[i] ?? 0) * 100) / (MTR[i] ?? 1));
  }

  const ADX = funMA(
    funArrDiv(funArrAbs(funArrSub(MDI, PDI)), funArrAdd(MDI, PDI)).map((v) =>
      v !== null ? v * 100 : null
    ),
    M
  );

  const ADXR = ADX.map((v, i) => {
    if (i < M) return v;
    if (ADX[i - M] === null) {
      return null;
    } else {
      return ((v ?? 0) + (ADX[i - M] ?? 0)) / 2;
    }
  });

  return {
    PDI: PDI.map((v) => funRound(v, 3)),
    MDI: MDI.map((v) => funRound(v, 3)),
    ADX: ADX.map((v) => funRound(v, 3)),
    ADXR: ADXR.map((v) => funRound(v, 3)),
  };
}
