package server

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


	"github.com/hashicorp/go-multierror"
	"golang.org/x/sync/errgroup"
)

// Server is an abstraction of a server that hosts multiple transports instances. Its responsibilities are to start
// the transports and allow them to gracefully shut down.
type Server interface {
	patterns.Runnable
}

// Transport is the interface for a transport instance that starts with the server. Unlike Server, it is always run
// in a context.
type Transport interface {
	StartWithContext(ctx context.Context) (err error)
	Stop() error
}

// serverImpl is an implementation of Server.
type serverImpl struct {
	patterns.Runnable
	config     *Config
	logger     logging.Logger
	transports []Transport
	quit       chan os.Signal
}

// NewServer returns a new instance of Server with given transports.
func NewServer(cp config.Provider, logger logging.Logger, transports ...Transport) Server {
	return &serverImpl{
		Runnable:   patterns.NewRunnable(),
		config:     GetConfig(cp),
		logger:     logger,
		transports: transports,
	}
}

// Start starts the server and blocks until it receives termination signals then gracefully stops.
func (s *serverImpl) Start() error {
	if err := s.Runnable.Start(); err != nil {
		return err
	}

	ctx := s.Context()
	mu := sync.Mutex{}
	var errs error

	for _, transport := range s.transports {
		t := transport

		go func() {
			if err := t.StartWithContext(ctx); err != nil {
				mu.Lock()
				defer mu.Unlock()

				errs = multierror.Append(errs, err)
			}
		}()
	}

	if s.config.StartUpTimeoutMs > 0 {
		// NOTE: there's a catch that the server cannot tell when its transports start-up routine would return.
		// This is due to the fact that most transport libraries (go-kit, gin, echo, etc.) would have a start
		// routine that either blocks or return error at an un-deterministic time.
		//
		// To work around, we give the server a start-up time out to wait, after which all its transports must
		// return error or will be assumed as running successfully.
		time.Sleep(time.Millisecond * time.Duration(s.config.StartUpTimeoutMs))
	}

	// There were errors, return.
	if errs != nil {
		return fmt.Errorf("failed to start transport with errors: %w", errs)
	}

	// We do not wait for all the
	s.logger.Info("server started successfully 🚀")

	s.quit = make(chan os.Signal)
	signal.Notify(s.quit, syscall.SIGINT, syscall.SIGTERM, syscall.SIGALRM)
	<-s.quit

	s.logger.Info("received terminated signal, server is shutting down ⏳")
	return s.Stop()
}

// Stop stops the server.
func (s *serverImpl) Stop() error {
	if err := s.shutdown(); err != nil {
		return err
	}

	return s.Runnable.Stop()
}

// shutdown shuts down the server and its transports.
func (s *serverImpl) shutdown() error {
	ctx, cancel := context.WithTimeout(s.Context(), time.Millisecond*time.Duration(s.config.ShutDownTimeoutMs))
	defer cancel()

	errChan := make(chan error)

	g, _ := errgroup.WithContext(ctx)
	for _, transport := range s.transports {
		t := transport
		g.Go(func() error {
			return t.Stop()
		})
	}

	go func() {
		errChan <- g.Wait()
	}()

	select {
	case <-ctx.Done():
		return errors.New("timeout waiting for transport to shutdown")
	case err := <-errChan:
		return err
	}
}
