package server

import (
	"context"
	"errors"
	"github.com/UpmeshLTD/urus-aio/infrastructure/config"
	"github.com/UpmeshLTD/urus-aio/infrastructure/logging"
	"os"
	"sync"
	"time"

	. "github.com/onsi/ginkgo/v2"
	. "github.com/onsi/gomega"
)

type testTransport struct {
	startFunc func(ctx context.Context) (err error)
	stopFunc  func() error
	startErr  error
	stopErr   error
}

func (t *testTransport) StartWithContext(ctx context.Context) (err error) {
	if t.startFunc != nil {
		return t.startFunc(ctx)
	}

	return t.startErr
}

func (t *testTransport) Stop() error {
	if t.stopFunc != nil {
		return t.stopFunc()
	}

	return t.stopErr
}

var _ = Describe("Server", func() {
	var (
		cp     config.Provider
		server Server
		trans  *testTransport
	)

	BeforeEach(func() {
		trans = &testTransport{}
		cp = config.NewProvider()
		server = NewServer(cp, logging.NewNopLogger(), trans)
	})

	Describe("Start", func() {
		It("should start with transports", func() {
			resultChan := make(chan bool)
			trans.startFunc = func(ctx context.Context) (err error) {
				resultChan <- true
				return nil
			}

			go func() {
				_ = server.Start()
			}()

			time.AfterFunc(time.Second, func() {
				resultChan <- false
			})

			result := <-resultChan
			Expect(result).Should(BeTrue())
		})

		It("should not start if transport fails", func() {
			trans.startErr = errors.New("some error")
			err := server.Start()
			Expect(err).ShouldNot(BeNil())
		})

		It("should not start if transport fails before start-up time out", func() {
			trans.startFunc = func(ctx context.Context) (err error) {
				time.Sleep(time.Millisecond * time.Duration(GetConfig(cp).StartUpTimeoutMs-100))
				return errors.New("some error")
			}
			err := server.Start()
			Expect(err).ShouldNot(BeNil())
		})

		It("should stop when receiving termination signals", func() {
			mu := sync.Mutex{}
			transportStopped := false
			trans.stopFunc = func() error {
				mu.Lock()
				defer mu.Unlock()

				transportStopped = true
				return nil
			}

			resultChan := make(chan bool)

			go func() {
				_ = server.Start()
				resultChan <- true
			}()

			time.Sleep(time.Millisecond * time.Duration(GetConfig(cp).StartUpTimeoutMs+100))
			impl := server.(*serverImpl)
			impl.quit <- os.Interrupt

			time.AfterFunc(time.Second, func() {
				resultChan <- false
			})

			result := <-resultChan
			Expect(result).Should(BeTrue())
			Expect(transportStopped).Should(BeTrue())
		})

		It("should return timeout if it takes too long to shut down", func() {
			mu := sync.Mutex{}
			transportStopped := false
			trans.stopFunc = func() error {
				mu.Lock()
				defer mu.Unlock()

				time.Sleep(time.Millisecond * 200)
				transportStopped = true
				return nil
			}

			resultChan := make(chan bool)

			go func() {
				_ = server.Start()
				resultChan <- true
			}()

			time.Sleep(time.Millisecond * time.Duration(GetConfig(cp).StartUpTimeoutMs+100))
			impl := server.(*serverImpl)
			impl.config.ShutDownTimeoutMs = 100
			impl.quit <- os.Interrupt

			time.AfterFunc(time.Second, func() {
				resultChan <- false
			})

			result := <-resultChan
			Expect(result).Should(BeTrue())
			Expect(transportStopped).Should(BeFalse())
		})
	})
})
