package io.delphiplatform.api.v3.bigtable.processing;

import org.junit.jupiter.api.Test;

import java.time.LocalDate;
import java.util.Collections;
import java.util.Comparator;
import java.util.List;
import java.util.Set;
import java.util.stream.Collectors;
import java.util.stream.Stream;

import io.delphiplatform.api.v3.constant.CountryCodeConstant;
import io.delphiplatform.api.v3.model.GroupByField;
import io.delphiplatform.api.v3.model.IncludeStreams;
import io.delphiplatform.api.v3.model.StreamModel;
import io.delphiplatform.api.v3.view.util.Params;

import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertTrue;

class StreamsAggregatorTest {

    private static final String US = "us";
    private static final String UK = "uk";
    private static final String PROJECT_NUMBER_1 = "GRAS_123";
    private static final String PROJECT_NUMBER_2 = "GRAS_456";
    private static final String PRODUCT_FAM_ID_1 = "GRAS_555";
    private static final String PRODUCT_FAM_ID_2 = "GRAS_777";
    private static final String DATE_1 = "2020-10-01";
    private static final String DATE_2 = "2020-10-02";

    private final StreamsAggregator aggregator = new StreamsAggregator(new StreamsMergingService());

    @Test
    void groupByAggregate_groupCountryCode() {
        String us = "us";
        String uk = "uk";
        List<StreamModel> grouped = aggregator.groupByAggregate(java.util.stream.Stream.of(
                createStream(10L, us),
                createStream(1L, us),
                createStream(30L, uk)
            ), Params.builder()
                .groupByFields(Set.of(GroupByField.COUNTRY))
                .build())
            .sorted(Comparator.comparing(StreamModel::getCountryCode))
            .collect(Collectors.toList());

        assertEquals(2, grouped.size());
        assertEquals(uk, grouped.get(0).getCountryCode());
        assertEquals(30, grouped.get(0).getStreams());
        assertEquals(us, grouped.get(1).getCountryCode());
        assertEquals(11, grouped.get(1).getStreams());
    }

    @Test
    void groupByAggregate_checkAppleAndAppleAgeBands() {
        List<StreamModel> grouped = aggregator.groupByAggregate(Stream.of(
                    createStream(10L, CountryCodeConstant.WORLDWIDE),
                    createStream(1L, CountryCodeConstant.WORLDWIDE)
                ),
                Params.builder()
                    .includeStreams(Set.of(IncludeStreams.DEMOGRAPHICS))
                    .groupByFields(Set.of(GroupByField.COUNTRY))
                    .build()
            )
            .collect(Collectors.toList());

        assertEquals(1, grouped.size());
        assertTrue(grouped.get(0).getAppleAgeBands().isPresent());
        assertTrue(grouped.get(0).getSpotifyAgeBands().isPresent());
        assertTrue(grouped.get(0).getGenders().isPresent());
        assertEquals(11, grouped.get(0).getStreams());
    }

    @Test
    void groupByAggregate_groupByCountryCodeAndDate() {
        Comparator<StreamModel> comparator = Comparator.comparing(StreamModel::getDate);
        List<StreamModel> grouped = aggregator.groupByAggregate(java.util.stream.Stream.of(
                    createStream(1L, US, DATE_1),
                    createStream(2L, UK, DATE_1),
                    createStream(100L, UK, DATE_1),
                    createStream(10L, US, DATE_2),
                    createStream(20L, UK, DATE_2)
                ),
                Params.builder()
                    .groupByFields(Set.of(GroupByField.DATE, GroupByField.COUNTRY))
                    .build()
            )
            .sorted(comparator.thenComparing(StreamModel::getCountryCode))
            .collect(Collectors.toList());

        assertEquals(4, grouped.size());

        assertEquals(DATE_1, grouped.get(0).getDate().toString());
        assertEquals(UK, grouped.get(0).getCountryCode());
        assertEquals(102, grouped.get(0).getStreams());

        assertEquals(DATE_1, grouped.get(1).getDate().toString());
        assertEquals(US, grouped.get(1).getCountryCode());
        assertEquals(1, grouped.get(1).getStreams());

        assertEquals(DATE_2, grouped.get(2).getDate().toString());
        assertEquals(UK, grouped.get(2).getCountryCode());
        assertEquals(20, grouped.get(2).getStreams());

        assertEquals(DATE_2, grouped.get(3).getDate().toString());
        assertEquals(US, grouped.get(3).getCountryCode());
        assertEquals(10, grouped.get(3).getStreams());
    }

    private StreamModel createStream(Long streams, String countryCode) {
        return createStream(streams, countryCode, null, null, null);
    }

    private StreamModel createStream(Long streams, String countryCode, String date) {
        return createStream(streams, countryCode, date, null, null);
    }

    private StreamModel createStream(Long streams, String countryCode, String date,
        String projectId, String productFamily) {
        StreamModel stream = new StreamModel();
        stream.date(date == null ? null : LocalDate.parse(date));
        stream.streams(streams);
        stream.countryCode(countryCode);
        stream.projectNumber(projectId);
        stream.productFamilyId(productFamily);
        return stream;
    }

    @Test
    void groupByAggregate_groupByProjectId() {
        Params params = Params.builder()
            .projectNumbers(Set.of(PROJECT_NUMBER_1, PROJECT_NUMBER_2))
            .groupByFields(Set.of(GroupByField.COUNTRY))
            .build();

        Comparator<StreamModel> comparator = Comparator.comparing(StreamModel::getProjectNumber)
            .thenComparing(StreamModel::getCountryCode);

        List<StreamModel> grouped = aggregator.groupByAggregate(Stream.of(
                createStream(1L, US, DATE_1, PROJECT_NUMBER_1, null),
                createStream(2L, UK, DATE_1, PROJECT_NUMBER_1, null),
                createStream(100L, UK, DATE_1, PROJECT_NUMBER_2, null),
                createStream(10L, US, DATE_2, PROJECT_NUMBER_1, null),
                createStream(20L, UK, DATE_2, PROJECT_NUMBER_2, null)
            ), params)
            .sorted(comparator)
            .collect(Collectors.toList());

        assertEquals(3, grouped.size());
        assertEquals(UK, grouped.get(0).getCountryCode());
        assertEquals(US, grouped.get(1).getCountryCode());
        assertEquals(UK, grouped.get(2).getCountryCode());
        assertEquals(2, grouped.get(0).getStreams());
        assertEquals(11, grouped.get(1).getStreams());
        assertEquals(120, grouped.get(2).getStreams());

    }

    @Test
    void groupByAggregate_groupByProductFamilyId() {
        Params params = Params.builder()
            .productFamilyIds(Set.of(PRODUCT_FAM_ID_1, PRODUCT_FAM_ID_2))
            .groupByFields(Set.of(GroupByField.COUNTRY))
            .build();

        Comparator<StreamModel> comparator = Comparator.comparing(StreamModel::getProductFamilyId)
            .thenComparing(StreamModel::getCountryCode);

        List<StreamModel> grouped = aggregator.groupByAggregate(Stream.of(
                createStream(1L, US, DATE_1, null, PRODUCT_FAM_ID_1),
                createStream(2L, UK, DATE_1, null, PRODUCT_FAM_ID_1),
                createStream(100L, UK, DATE_1, null, PRODUCT_FAM_ID_2),
                createStream(10L, US, DATE_2, null, PRODUCT_FAM_ID_2),
                createStream(20L, UK, DATE_2, null, PRODUCT_FAM_ID_1)
            ), params)
            .sorted(comparator)
            .collect(Collectors.toList());

        assertEquals(4, grouped.size());
        assertEquals(UK, grouped.get(0).getCountryCode());
        assertEquals(US, grouped.get(1).getCountryCode());
        assertEquals(UK, grouped.get(2).getCountryCode());
        assertEquals(US, grouped.get(3).getCountryCode());
        assertEquals(22, grouped.get(0).getStreams());
        assertEquals(1, grouped.get(1).getStreams());
        assertEquals(100, grouped.get(2).getStreams());
        assertEquals(10, grouped.get(3).getStreams());

    }

}
