import type {
  DefaultError,
  InvalidateQueryFilters,
  QueryClient,
  QueryKey,
  UseMutationOptions,
} from '@tanstack/react-query'
import type { QueryBroadcaster } from './broadcast'

export interface OptimisticMutationContext<TQueryData> {
  snapshot: TQueryData | undefined
  rollback: () => void
}

export interface CreateOptimisticMutationOptions<
  TData,
  TError = DefaultError,
  TVariables = void,
  TQueryData = unknown,
> {
  queryClient: QueryClient
  queryKey: QueryKey
  mutationFn: (variables: TVariables) => Promise<TData>
  apply: (current: TQueryData | undefined, variables: TVariables) => TQueryData | undefined
  invalidate?: false | QueryKey | InvalidateQueryFilters
  broadcaster?: Pick<QueryBroadcaster, 'publishKeys'>
  onMutate?: (
    variables: TVariables,
  ) => Promise<void | OptimisticMutationContext<TQueryData>> | void | OptimisticMutationContext<TQueryData>
  onError?: (
    error: TError,
    variables: TVariables,
    context: OptimisticMutationContext<TQueryData> | undefined,
  ) => Promise<unknown> | unknown
  onSuccess?: (
    data: TData,
    variables: TVariables,
    context: OptimisticMutationContext<TQueryData>,
  ) => Promise<unknown> | unknown
  onSettled?: (
    data: TData | undefined,
    error: TError | null,
    variables: TVariables,
    context: OptimisticMutationContext<TQueryData> | undefined,
  ) => Promise<unknown> | unknown
}

export type OptimisticMutationOptions<
  TData,
  TError = DefaultError,
  TVariables = void,
  TQueryData = unknown,
> = Omit<
  UseMutationOptions<TData, TError, TVariables, OptimisticMutationContext<TQueryData>>,
  'mutationFn' | 'onMutate' | 'onError' | 'onSuccess' | 'onSettled'
> & {
  mutationFn: (variables: TVariables) => Promise<TData>
  onMutate: (variables: TVariables) => Promise<OptimisticMutationContext<TQueryData>>
  onError: (
    error: TError,
    variables: TVariables,
    context: OptimisticMutationContext<TQueryData> | undefined,
  ) => Promise<void>
  onSuccess: (
    data: TData,
    variables: TVariables,
    context: OptimisticMutationContext<TQueryData> | undefined,
  ) => Promise<void>
  onSettled: (
    data: TData | undefined,
    error: TError | null,
    variables: TVariables,
    context: OptimisticMutationContext<TQueryData> | undefined,
  ) => Promise<void>
}

function toInvalidateFilters(
  queryKey: QueryKey,
  invalidate: false | QueryKey | InvalidateQueryFilters | undefined,
): false | InvalidateQueryFilters {
  if (invalidate === false) return false
  if (!invalidate) return { queryKey }

  const filters: InvalidateQueryFilters = Array.isArray(invalidate)
    ? { queryKey: invalidate }
    : (invalidate as InvalidateQueryFilters)

  return filters
}

export function createOptimisticMutation<
  TData,
  TError = DefaultError,
  TVariables = void,
  TQueryData = unknown,
>({
  queryClient,
  queryKey,
  mutationFn,
  apply,
  invalidate,
  broadcaster,
  onMutate,
  onError,
  onSuccess,
  onSettled,
}: CreateOptimisticMutationOptions<TData, TError, TVariables, TQueryData>): UseMutationOptions<
  TData,
  TError,
  TVariables,
  OptimisticMutationContext<TQueryData>
> &
  OptimisticMutationOptions<TData, TError, TVariables, TQueryData> {
  return {
    mutationFn,
    async onMutate(variables) {
      await queryClient.cancelQueries({ queryKey })

      const snapshot = queryClient.getQueryData<TQueryData>(queryKey)
      const next = apply(snapshot, variables)

      queryClient.setQueryData(queryKey, next)
      broadcaster?.publishKeys([queryKey])

      const context: OptimisticMutationContext<TQueryData> = {
        snapshot,
        rollback: () => {
          queryClient.setQueryData(queryKey, snapshot)
          broadcaster?.publishKeys([queryKey])
        },
      }

      const extraContext = await onMutate?.(variables)
      return extraContext ?? context
    },
    async onError(error, variables, context) {
      context?.rollback()
      await onError?.(error, variables, context)
    },
    async onSuccess(data, variables, context) {
      if (context) {
        await onSuccess?.(data, variables, context)
      }
    },
    async onSettled(data, error, variables, context) {
      const filters = toInvalidateFilters(queryKey, invalidate)
      if (filters !== false) {
        await queryClient.invalidateQueries(filters)
        broadcaster?.publishKeys([filters.queryKey ?? queryKey])
      }

      await onSettled?.(data, error ?? null, variables, context)
    },
  }
}
