import java.util.*;

class Main {
    public static void main(String[] args) {
        int[] nums = {1, 2, 1, 2, 3};
        int k = 2;

        if (k <= 0) {
            System.out.println(0);
            return;
        }

        int n = nums.length;
        long count = 0;
        int left = 0;
        Map<Integer, Integer> freqMap = new HashMap<>();

        for (int right = 0; right < n; right++) {
            freqMap.put(nums[right], freqMap.getOrDefault(nums[right], 0) + 1);

            while (freqMap.size() > k) {
                int leftElement = nums[left];
                freqMap.put(leftElement, freqMap.get(leftElement) - 1);
                if (freqMap.get(leftElement) == 0) {
                    freqMap.remove(leftElement);
                }
                left++;
            }

            count += (right - left + 1);
        }

        System.out.println(count);
    }
}
