package main

import (
	"errors"
	"io"
	"net"
	"sync"
	"testing"
	"time"
)

var (
	errTodo     = errors.New("The code you have to implements")
	errConflict = errors.New("conflict with other connections")
)

// A new net.Conn which compile multiple net.Conn
// to speed up transmition
type groupConn struct {
	mu           sync.Mutex
	main         net.Conn
	standByConns []net.Conn
	isClient     bool
}

func NewServerConn(main net.Conn, n int) (*groupConn, error) {
	c, err := newServerGroupConn(main, n)
	return c, err
}

func NewClientConn(main net.Conn, n int) (*groupConn, error) {
	c, err := newClientGroupConn(main, n)
	return c, err
}

func newClientGroupConn(main net.Conn, n int) (*groupConn, error) {
	gc := &groupConn{
		main:         main,
		standByConns: make([]net.Conn, n+1),
		isClient:     true,
	}
	gc.standByConns[0] = main
	return gc, nil
}

func newServerGroupConn(main net.Conn, n int) (*groupConn, error) {
	gc := &groupConn{
		main:         main,
		standByConns: make([]net.Conn, n+1),
		isClient:     true,
	}
	gc.standByConns[0] = main
	return gc, nil
}

func (gc *groupConn) putStandByConn(c net.Conn, i int) error {
	gc.mu.Lock()
	defer gc.mu.Unlock()
	if gc.standByConns[i+1] != nil {
		return errConflict
	}
	gc.standByConns[i+1] = c
	return nil
}

func (gc *groupConn) Write(b []byte) (n int, err error) {
	// TODO
	err = errTodo
	return
}

func (gc *groupConn) Read(b []byte) (n int, err error) {
	// TODO
	err = errTodo
	return
}

func (gc *groupConn) Close() error {
	// TODO
	return errTodo
}

func (gc *groupConn) LocalAddr() net.Addr {
	return gc.main.LocalAddr()
}

func (gc *groupConn) RemoteAddr() net.Addr {
	return gc.main.RemoteAddr()
}

func (gc *groupConn) SetDeadline(t time.Time) error {
	// TODO
	return errTodo
}

func (gc *groupConn) SetReadDeadline(t time.Time) error {
	// TODO
	return errTodo
}

func (gc *groupConn) SetWriteDeadline(t time.Time) error {
	// TODO
	return errTodo
}

func setupServerTest(t *testing.T, handleConnectionFn func(net.Conn)) (addr string, stopfunc func(), client net.Conn, err error) {
	ln, err := net.Listen("tcp", "localhost:0")
	if err != nil {
		return "", nil, nil, err
	}
	go func() {
		conn, err := ln.Accept()
		if err != nil {
			t.Error(err)
			return
		}
		go handleConnectionFn(conn)
	}()
	time.Sleep(20 * time.Millisecond)
	addr = ln.Addr().String()
	conn, err := net.Dial("tcp", addr)
	if err != nil {
		ln.Close()
		return "", nil, nil, err
	}
	return ln.Addr().String(), func() { ln.Close() }, conn, nil
}

func TestGroupConnNormalTransfer(t *testing.T) {
	const connSize = 9
	var serverGroupConn *groupConn

	getMainConn := func(c net.Conn) {
		var err error
		serverGroupConn, err = newServerGroupConn(c, connSize)
		if err != nil {
			t.Fatal(err)
		}
	}
	_, stopfunc, c1, err := setupServerTest(t, getMainConn)
	if err != nil {
		t.Fatal(err)
	}
	defer stopfunc()
	clientGroupConn, err := newClientGroupConn(c1, connSize)
	if err != nil {
		t.Fatal(err)
	}

	for i := 0; i < connSize; i++ {
		go func(idx int) {
			_, _, c2, err := setupServerTest(t, func(c3 net.Conn) {
				serverGroupConn.putStandByConn(c3, idx)
			})
			if err != nil {
				t.Fatal(err)
			}
			clientGroupConn.putStandByConn(c2, idx)
		}(i)
	}

	//TODO random data
	testBuf1 := []byte("test ajxjalkaj;dkasdf")
	testBuf2 := []byte("test a;lk34j1j;vxjalkaj;dkasdf asdfasdb")
	testBufBigData := []byte("test a;lk34j1j;vxjalkaj;dkasdf asdfasdb")
	go func() {
		buf2 := make([]byte, 4096)
		n, err := io.ReadFull(serverGroupConn, buf2[:len(testBuf1)])
		if err != nil {
			t.Fatal(err)
		}
		if n != len(testBuf1) {
			t.Fatal("error size")
		}

		n, err = serverGroupConn.Write(testBuf2)
		if err != nil {
			t.Fatal(err)
		}

		if n != len(testBuf2) {
			t.Fatal("error size")
		}
	}()

	// client side
	{
		n, err := clientGroupConn.Write(testBuf1)
		if err != nil {
			t.Fatal(err)
		}

		if n != len(testBuf1) {
			t.Fatal("error size")
		}

		buf3 := make([]byte, 4096)
		n, err = io.ReadFull(clientGroupConn, buf3[:len(testBuf2)])
		if err != nil {
			t.Fatal(err)
		}
		if n != len(testBuf2) {
			t.Fatal("error size")
		}

		n, err = clientGroupConn.Write(testBufBigData)
		if err != nil {
			t.Fatal(err)
		}

		if n != len(testBufBigData) {
			t.Fatal("error size")
		}
	}

}

func TestGroupConnHighSendSlowRecvTransfer(t *testing.T) {
	// TODO
}

func TestGroupConnLowSendHighRecvTransfer(t *testing.T) {
	// TODO
}

func TestGroupConnDifferentRatesTransfer(t *testing.T) {
	// TODO
}

func TestGroupConnClose(t *testing.T) {
	// TODO
}

func TestGroupConnWriteDeadline(t *testing.T) {
	// TODO
}

func TestGroupConnReadDeadline(t *testing.T) {
	// TODO
}

func TestGroupConnRecvInOrder(t *testing.T) {
	// TODO
}

func TestGroupConnResendDroppedSequence(t *testing.T) {
	// TODO
}
