// Port of sqlglot/typing/snowflake.py.

///|
let snowflake_date_parts : Array[String] = [
  "DAY", "WEEK", "MONTH", "QUARTER", "YEAR",
]

///|
const SNOWFLAKE_MAX_PRECISION : Int = 38

///|
const SNOWFLAKE_MAX_SCALE : Int = 37

///|
fn snowflake_datatype(
  sql : String,
  s : TypeAnnotator,
) -> @core.Expr raise @core.SqlglotError {
  datatype_from_str_in(sql, "snowflake", s.dialect)
}

///|
fn snowflake_annotate_reverse(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit {
  s.annotate_by_args(expression, [Key("this")])
  if expression.is_type([NULL]) {
    // Snowflake treats REVERSE(NULL) as a VARCHAR
    s.set_dtype(expression, VARCHAR)
  }
}

///|
/// TIMESTAMP_FROM_PARTS with time_zone -> TIMESTAMPTZ, otherwise TIMESTAMP (NTZ).
fn snowflake_annotate_timestamp_from_parts(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit {
  s.set_dtype(
    expression,
    if expression.has("zone") {
      TIMESTAMPTZ
    } else {
      TIMESTAMP
    },
  )
}

///|
fn snowflake_annotate_date_or_time_add(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit {
  if opt_is_type(expression.this(), [DATE]) &&
    !snowflake_date_parts.contains(@core.py_upper(expression.text("unit"))) {
    s.set_dtype(expression, TIMESTAMPNTZ)
  } else {
    s.annotate_by_args(expression, [Key("this")])
  }
}

///|
/// DECODE(expr, val1, ret1, val2, ret2, ..., default): the type is inferred from the
/// return values only.
fn snowflake_annotate_decode_case(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit {
  let expressions = expression.expressions()
  let return_types : Array[TType?] = []
  for i = 2; i < expressions.length(); i = i + 2 {
    return_types.push(type_of(expressions[i]))
  }
  // If the total number of expressions is even, the last one is the default
  if expressions.length() % 2 == 0 && !expressions.is_empty() {
    return_types.push(type_of(expressions[expressions.length() - 1]))
  }
  let mut last_type : TType? = None
  for ret_type in return_types {
    let t1 = match last_type {
      Some(_) => last_type
      None => ret_type
    }
    last_type = s.maybe_coerce_opt(t1, ret_type)
  }
  s.set_type(expression, last_type)
}

///|
fn snowflake_annotate_arg_max_min(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit {
  if expression.has("count") {
    s.set_dtype(expression, ARRAY)
  } else {
    s.set_type_of(expression, expression.this())
  }
}

///|
/// For PERCENTILE_DISC / PERCENTILE_CONT, the type is the ordered expression's type.
fn snowflake_annotate_within_group(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit {
  let this = expression.this()
  let ordered_this = match (this, expression.expression()) {
    (Some(t), Some(order)) if t.kind.is_any([PercentileDisc, PercentileCont]) &&
      order.kind.is_a(Order) &&
      order.expressions().length() == 1 &&
      order.expressions()[0].kind.is_a(Ordered) =>
      Some(order.expressions()[0].this())
    _ => None
  }
  match ordered_this {
    Some(ot) => s.set_type_of(expression, ot)
    None => s.set_type_of(expression, this)
  }
}

///|
/// MEDIAN: FLOAT/DOUBLE -> DOUBLE, NUMBER(p, s) -> NUMBER(min(p+3, 38), min(s+3, 37)).
fn snowflake_annotate_median(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit raise @core.SqlglotError {
  s.annotate_by_args(expression, [Key("this")])
  let input_type = expression.this().bind(e => e.get_type())
  if opt_is_type(input_type, [DOUBLE]) {
    s.set_dtype(expression, DOUBLE)
  } else {
    let exprs = input_type.map(t => t.expressions()).unwrap_or([])
    let precision = datatype_param_int(exprs.get(0), SNOWFLAKE_MAX_PRECISION)
    let scale = datatype_param_int(exprs.get(1), 0)
    let new_precision = @core.min_int(precision + 3, SNOWFLAKE_MAX_PRECISION)
    let new_scale = @core.min_int(scale + 3, SNOWFLAKE_MAX_SCALE)
    s.set_type_expr(
      expression,
      Some(snowflake_datatype("NUMBER(\{new_precision}, \{new_scale})", s)),
    )
  }
}

///|
/// VAR_POP, VAR_SAMP, VARIANCE, VARIANCE_POP: DECFLOAT -> DECFLOAT(38),
/// FLOAT/DOUBLE -> DOUBLE, INT/NUMBER(p, 0) -> NUMBER(38, 6), NUMBER(p, s) -> NUMBER(38, max(12, s)).
fn snowflake_annotate_variance(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit raise @core.SqlglotError {
  s.annotate_by_args(expression, [Key("this")])
  let input_type = expression.this().bind(e => e.get_type())
  if opt_is_type(input_type, [DECFLOAT]) {
    s.set_type_expr(expression, Some(snowflake_datatype("DECFLOAT", s)))
  } else if opt_is_type(input_type, [FLOAT, DOUBLE]) {
    s.set_dtype(expression, DOUBLE)
  } else {
    let exprs = input_type.map(t => t.expressions()).unwrap_or([])
    let scale = datatype_param_int(exprs.get(1), 0)
    let new_scale = if scale == 0 { 6 } else if scale > 12 { scale } else { 12 }
    s.set_type_expr(
      expression,
      Some(
        snowflake_datatype(
          "NUMBER(\{SNOWFLAKE_MAX_PRECISION}, \{new_scale})",
          s,
        ),
      ),
    )
  }
}

///|
/// KURTOSIS: DECFLOAT -> DECFLOAT, DOUBLE/FLOAT -> DOUBLE, other numerics -> NUMBER(38, 12).
fn snowflake_annotate_kurtosis(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit raise @core.SqlglotError {
  s.annotate_by_args(expression, [Key("this")])
  let input_type = expression.this().bind(e => e.get_type())
  if opt_is_type(input_type, [DECFLOAT]) {
    s.set_type_expr(expression, Some(snowflake_datatype("DECFLOAT", s)))
  } else if opt_is_type(input_type, [FLOAT, DOUBLE]) {
    s.set_dtype(expression, DOUBLE)
  } else {
    s.set_type_expr(
      expression,
      Some(snowflake_datatype("NUMBER(\{SNOWFLAKE_MAX_PRECISION}, 12)", s)),
    )
  }
}

///|
/// Math functions that preserve DECFLOAT but return DOUBLE for other inputs.
fn snowflake_annotate_math_with_float_decfloat(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit {
  s.annotate_by_args(expression, [Key("this")])
  if opt_is_type(expression.this(), [DECFLOAT]) {
    s.set_type_of(expression, expression.this())
  } else {
    s.set_dtype(expression, DOUBLE)
  }
}

///|
fn snowflake_annotate_str_to_time(
  s : TypeAnnotator,
  expression : @core.Expr,
) -> Unit {
  // target_type is stored as a DataType instance
  match expression.arg("target_type") {
    Some(t) if t.kind.is_a(DataType) =>
      match t.datatype_this() {
        Some(d) => s.set_dtype(expression, d)
        None => s.set_type(expression, None)
      }
    _ => s.set_dtype(expression, TIMESTAMP)
  }
}

///|
/// Python `sqlglot.typing.snowflake.EXPRESSION_METADATA`.
fn snowflake_expression_metadata() -> ExprMetadata {
  let m = extend_metadata(base_expression_metadata)
  annotate_all(
    m,
    [
      AddMonths,
      Ceil,
      DateTrunc,
      Floor,
      Left,
      Mode,
      Pad,
      Right,
      Round,
      Stuff,
      Substring,
      TimeSlice,
      TimestampTrunc,
    ],
    (s, e) => s.annotate_by_args(e, [Key("this")]),
  )
  returns_all(
    m,
    [
      ApproxTopK,
      ApproxTopKEstimate,
      Array,
      ArrayAgg,
      ArrayAppend,
      ArrayCompact,
      ArrayConcat,
      ArrayConstructCompact,
      ArrayPrepend,
      ArrayRemove,
      ArraysZip,
      ArrayUniqueAgg,
      ArrayUnionAgg,
      MapKeys,
      RegexpExtractAll,
      Split,
      StringToArray,
      StrtokToArray,
    ],
    ARRAY,
  )
  returns_all(
    m,
    [
      BitmapBitPosition,
      BitmapBucketNumber,
      BitmapCount,
      Factorial,
      GroupingId,
      MD5NumberLower64,
      MD5NumberUpper64,
      Rand,
      Seq8,
      Zipf,
    ],
    BIGINT,
  )
  returns_all(
    m,
    [
      Base64DecodeBinary,
      BitmapConstructAgg,
      BitmapOrAgg,
      Compress,
      DecompressBinary,
      Decrypt,
      DecryptRaw,
      Encrypt,
      EncryptRaw,
      HexString,
      MD5Digest,
      SHA1Digest,
      SHA2Digest,
      ToBinary,
      TryBase64DecodeBinary,
      TryHexDecodeBinary,
      Unhex,
    ],
    BINARY,
  )
  returns_all(
    m,
    [
      Booland,
      Boolnot,
      Boolor,
      BoolxorAgg,
      EqualNull,
      IsNullValue,
      MapContainsKey,
      Search,
      SearchIp,
      ToBoolean,
    ],
    BOOLEAN,
  )
  returns_all(m, [NextDay, PreviousDay], DATE)
  annotate_all(
    m,
    [
      BitwiseAndAgg,
      BitwiseOrAgg,
      BitwiseXorAgg,
      RegexpCount,
      RegexpInstr,
      ToNumber,
    ],
    (s, e) => s.set_type_expr(e, Some(snowflake_datatype("NUMBER", s))),
  )
  returns_all(
    m,
    [
      ApproxPercentileEstimate,
      ApproximateSimilarity,
      CosineDistance,
      DotProduct,
      EuclideanDistance,
      ManhattanDistance,
      MonthsBetween,
      Normal,
    ],
    DOUBLE,
  )
  m[Kurtosis] = Annotator((s, e) => snowflake_annotate_kurtosis(s, e))
  returns_all(m, [ToDecfloat, TryToDecfloat], DECFLOAT)
  annotate_all(
    m,
    [
      Acos,
      Asin,
      Atan,
      Atan2,
      Cbrt,
      Cos,
      Cot,
      Degrees,
      Exp,
      Ln,
      Log,
      Pow,
      Radians,
      RegrAvgx,
      RegrAvgy,
      RegrCount,
      RegrIntercept,
      RegrR2,
      RegrSlope,
      RegrSxx,
      RegrSxy,
      RegrSyy,
      RegrValx,
      RegrValy,
      Sin,
      Sqrt,
      Tan,
      Tanh,
    ],
    (s, e) => snowflake_annotate_math_with_float_decfloat(s, e),
  )
  returns_all(
    m,
    [
      ByteLength,
      DenseRank,
      Grouping,
      JarowinklerSimilarity,
      MapSize,
      Minute,
      Ntile,
      Rank,
      RowNumber,
      RtrimmedLength,
      Second,
      Seq1,
      Seq2,
      Seq4,
      WidthBucket,
    ],
    INT,
  )
  returns_all(
    m,
    [
      ApproxPercentileAccumulate,
      ApproxPercentileCombine,
      ApproxTopKAccumulate,
      ApproxTopKCombine,
      ObjectAgg,
      ParseIp,
      ParseUrl,
      XMLGet,
    ],
    OBJECT,
  )
  returns_all(m, [MapCat, MapDelete, MapInsert, MapPick], MAP)
  returns_all(m, [ToFile], FILE)
  returns_all(m, [TimeFromParts, TsOrDsToTime], TIME)
  returns_all(m, [CurrentTimestamp, Localtimestamp], TIMESTAMPLTZ)
  returns_all(m, [DayOfMonth, DayOfWeek, DayOfYear, Quarter], TINYINT)
  returns_all(
    m,
    [
      AIAgg,
      AIClassify,
      AISummarizeAgg,
      Base64DecodeString,
      Base64Encode,
      CheckJson,
      CheckXml,
      Collate,
      Collation,
      CurrentAccount,
      CurrentAccountName,
      CurrentAvailableRoles,
      CurrentClient,
      CurrentDatabase,
      CurrentIpAddress,
      CurrentSchemas,
      CurrentSecondaryRoles,
      CurrentSession,
      CurrentStatement,
      CurrentTransaction,
      CurrentWarehouse,
      CurrentOrganizationUser,
      CurrentRegion,
      CurrentRoleType,
      CurrentOrganizationName,
      DecompressString,
      HexDecodeString,
      Hex,
      Randstr,
      RegexpExtract,
      RegexpReplace,
      Replace,
      Soundex,
      SoundexP123,
      SplitPart,
      Strtok,
      TryBase64DecodeString,
      TryHexDecodeString,
      Uuid,
    ],
    VARCHAR,
  )
  returns_all(m, [Minhash, MinhashCombine], VARIANT)
  annotate_all(m, [Variance, VariancePop], (s, e) => {
    snowflake_annotate_variance(s, e)
  })
  m[ArgMax] = Annotator((s, e) => snowflake_annotate_arg_max_min(s, e))
  m[ArgMin] = Annotator((s, e) => snowflake_annotate_arg_max_min(s, e))
  m[ConcatWs] = Annotator((s, e) => s.annotate_by_args(e, [Key("expressions")]))
  m[ConvertTimezone] = Annotator((s, e) => {
    s.set_dtype(e, if e.has("source_tz") { TIMESTAMPNTZ } else { TIMESTAMPTZ })
  })
  m[DateAdd] = Annotator((s, e) => snowflake_annotate_date_or_time_add(s, e))
  m[DecodeCase] = Annotator((s, e) => snowflake_annotate_decode_case(s, e))
  m[HashAgg] = Annotator((s, e) => {
    s.set_type_expr(e, Some(snowflake_datatype("NUMBER(19, 0)", s)))
  })
  m[Median] = Annotator((s, e) => snowflake_annotate_median(s, e))
  m[Reverse] = Annotator((s, e) => snowflake_annotate_reverse(s, e))
  m[StrToTime] = Annotator((s, e) => snowflake_annotate_str_to_time(s, e))
  m[TimeAdd] = Annotator((s, e) => snowflake_annotate_date_or_time_add(s, e))
  m[TimestampFromParts] = Annotator((s, e) => {
    snowflake_annotate_timestamp_from_parts(s, e)
  })
  m[WithinGroup] = Annotator((s, e) => snowflake_annotate_within_group(s, e))
  m
}