package io.delphiplatform.api.v3.bigtable;

import com.google.api.gax.rpc.ServerStream;
import com.google.cloud.bigtable.data.v2.BigtableDataClient;
import com.google.cloud.bigtable.data.v2.models.Query;
import com.google.cloud.bigtable.data.v2.models.Row;
import com.google.common.collect.Lists;

import org.apache.commons.lang3.tuple.Pair;
import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.mockito.ArgumentCaptor;
import org.mockito.Mockito;

import java.util.Comparator;
import java.util.List;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ExecutionException;
import java.util.concurrent.Executors;
import java.util.stream.Collectors;
import java.util.stream.IntStream;
import java.util.stream.Stream;

import io.delphiplatform.api.v3.bigtable.client.DelphiBigtableClient;
import io.delphiplatform.api.v3.bigtable.config.RequestRowKeyHandle;
import io.delphiplatform.api.v3.bigtable.config.RowKeyRequestType;

import static io.delphiplatform.api.utils.TestUtils.getRanges;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;

class BigtableQueryExecutorTest {

    private BigtableQueryExecutor queryExecutor;

    private BigtableDataClient btClient = mock(BigtableDataClient.class);
    private DelphiBigtableClient client = new DelphiBigtableClient(btClient);

    @BeforeEach
    public void setUp() {
        queryExecutor = new BigtableQueryExecutor(client, Executors.newFixedThreadPool(16), 100);
    }

    @Test
    void executeByRanges_joinsAllRows() throws ExecutionException, InterruptedException {
        Row row1 = mock(Row.class);
        Row row2 = mock(Row.class);

        ServerStream<Row> serverStream1 = createClientResponse(row1);
        ServerStream<Row> serverStream2 = createClientResponse(row2);

        ArgumentCaptor<Query> argumentCaptor = ArgumentCaptor.forClass(Query.class);
        Mockito.when(btClient.readRows(argumentCaptor.capture())).thenReturn(serverStream1, serverStream2);

        CompletableFuture<Stream<Row>> allRowsFuture =
            queryExecutor.executePartitioned("testTable", createRangesList(150));

        List<Row> allRows = allRowsFuture.get().collect(Collectors.toList());

        Assertions.assertTrue(allRows.contains(row1));
        Assertions.assertTrue(allRows.contains(row2));

        List<Query> allQueries = argumentCaptor.getAllValues();
        allQueries.sort(Comparator.comparing(q -> getRanges(q).size()));
        Assertions.assertEquals(2, allQueries.size());
        Assertions.assertEquals(50, getRanges(allQueries.get(0)).size());
        Assertions.assertEquals(100, getRanges(allQueries.get(1)).size());
    }

    private ServerStream<Row> createClientResponse(Row row2) {
        ServerStream<Row> serverStream2 = mock(ServerStream.class);
        when(serverStream2.spliterator()).thenReturn(Lists.newArrayList(row2).spliterator());
        return serverStream2;
    }

    private List<RequestRowKeyHandle> createRangesList(int count) {
        return IntStream.range(0, count)
            .mapToObj(i -> RequestRowKeyHandle.builder()
                .rowKeyPrefix(null)
                .rowKeyRequestType(RowKeyRequestType.RANGE)
                .rowKey(Pair.of(i + "start", i + "end"))
                .build())
            .collect(Collectors.toList());
    }
}
