// Package cloudwatchwriter
// https://github.com/mec07/cloudwatchwriter
// https://github.com/kdar/logrus-cloudwatchlogs
// https://github.com/rs/zerolog
package cloudwatchwriter

import (
	"sync"
	"time"

	"github.com/aws/aws-sdk-go/aws"
	"github.com/aws/aws-sdk-go/aws/awserr"
	"github.com/aws/aws-sdk-go/service/cloudwatchlogs"
	"github.com/pkg/errors"
)

const (
	// minBatchInterval is 200 ms as the maximum rate of PutLogEvents is 5
	// requests per second.
	minBatchInterval time.Duration = 200000000
	// batchSizeLimit is 1MB in bytes, the limit imposed by AWS CloudWatch Logs
	// on the size the batch of logs we send, see:
	// https://docs.aws.amazon.com/AmazonCloudWatchLogs/latest/APIReference/API_PutLogEvents.html
	batchSizeLimit = 1048576
	// maxNumLogEvents is the maximum number of messages that can be sent in one
	// batch, also an AWS limitation, see:
	// https://docs.aws.amazon.com/AmazonCloudWatchLogs/latest/APIReference/API_PutLogEvents.html
	maxNumLogEvents = 10000
	// additionalBytesPerLogEvent is the number of additional bytes per log
	// event, other than the length of the log message, see:
	// https://docs.aws.amazon.com/AmazonCloudWatchLogs/latest/APIReference/API_PutLogEvents.html
	additionalBytesPerLogEvent = 26
)

// CloudWatchLogsClient represents the AWS cloudwatchlogs client that we need to talk to CloudWatch
type CloudWatchLogsClient interface {
	DescribeLogStreams(*cloudwatchlogs.DescribeLogStreamsInput) (*cloudwatchlogs.DescribeLogStreamsOutput, error)
	CreateLogGroup(*cloudwatchlogs.CreateLogGroupInput) (*cloudwatchlogs.CreateLogGroupOutput, error)
	CreateLogStream(*cloudwatchlogs.CreateLogStreamInput) (*cloudwatchlogs.CreateLogStreamOutput, error)
	PutLogEvents(*cloudwatchlogs.PutLogEventsInput) (*cloudwatchlogs.PutLogEventsOutput, error)
}

// CloudWatchWriter can be inserted into zerolog to send logs to CloudWatch.
type CloudWatchWriter struct {
	client            CloudWatchLogsClient
	batchInterval     time.Duration
	logGroupName      *string
	logStreamName     *string
	nextSequenceToken *string
	err               *error
	ch                chan *cloudwatchlogs.InputLogEvent
	m                 sync.Mutex
}

// New returns a pointer to a CloudWatchWriter struct, or an error.
func New(client CloudWatchLogsClient, batchInterval time.Duration, logGroupName, logStreamName string) (*CloudWatchWriter, error) {
	if batchInterval < minBatchInterval {
		return nil, errors.New("supplied batch interval is less than the minimum")
	}

	writer := &CloudWatchWriter{
		client:        client,
		batchInterval: batchInterval,
		logGroupName:  aws.String(logGroupName),
		logStreamName: aws.String(logStreamName),
		ch:            make(chan *cloudwatchlogs.InputLogEvent, maxNumLogEvents),
	}

	logStream, err := writer.getOrCreateLogStream()
	if err != nil {
		return nil, err
	}

	writer.nextSequenceToken = logStream.UploadSequenceToken

	go writer.putBatches()

	return writer, nil
}

// Write implements the io.Writer interface.
func (c *CloudWatchWriter) Write(log []byte) (int, error) {
	event := &cloudwatchlogs.InputLogEvent{
		Message: aws.String(string(log)),
		// Timestamp has to be in milliseconds since the epoch
		Timestamp: aws.Int64(time.Now().UTC().UnixNano() / int64(time.Millisecond)),
	}

	c.ch <- event

	if c.err != nil {
		lastErr := *c.err
		c.err = nil
		return 0, lastErr
	}

	return len(log), nil
}

func (c *CloudWatchWriter) putBatches() {
	var batch []*cloudwatchlogs.InputLogEvent
	ticker := time.Tick(c.batchInterval)
	size := 0
	for {
		select {
		case p := <-c.ch:
			messageSize := len(*p.Message) + additionalBytesPerLogEvent
			if size+messageSize >= batchSizeLimit || len(batch) == maxNumLogEvents {
				go c.sendBatch(batch, 0)
				batch = nil
				size = 0
			}
			batch = append(batch, p)
			size += messageSize
		case <-ticker:
			go c.sendBatch(batch, 0)
			batch = nil
			size = 0
		}
	}
}

// Only allow 1 retry of an invalid sequence token.
func (c *CloudWatchWriter) sendBatch(batch []*cloudwatchlogs.InputLogEvent, retryNum int) {
	c.m.Lock()
	defer c.m.Unlock()

	if retryNum > 1 || len(batch) == 0 {
		return
	}

	input := &cloudwatchlogs.PutLogEventsInput{
		LogEvents:     batch,
		LogGroupName:  c.logGroupName,
		LogStreamName: c.logStreamName,
		SequenceToken: c.nextSequenceToken,
	}

	output, err := c.client.PutLogEvents(input)
	if err != nil {
		if invalidSequenceTokenErr, ok := err.(*cloudwatchlogs.InvalidSequenceTokenException); ok {
			c.nextSequenceToken = invalidSequenceTokenErr.ExpectedSequenceToken
			go c.sendBatch(batch, retryNum+1)
			return
		}
		c.err = &err
		return
	}
	c.nextSequenceToken = output.NextSequenceToken
}

// getOrCreateLogStream gets info on the log stream for the log group and log
// stream we're interested in -- primarily for the purpose of finding the value
// of the next sequence token. If the log group doesn't exist, then we create
// it, if the log stream doesn't exist, then we create it.
func (c *CloudWatchWriter) getOrCreateLogStream() (*cloudwatchlogs.LogStream, error) {
	// Get the log streams that match our log group name and log stream
	output, err := c.client.DescribeLogStreams(&cloudwatchlogs.DescribeLogStreamsInput{
		LogGroupName:        c.logGroupName,
		LogStreamNamePrefix: c.logStreamName,
	})
	if err != nil || output == nil {
		awserror, ok := err.(awserr.Error)
		// i.e. the log group does not exist
		if ok && awserror.Code() == cloudwatchlogs.ErrCodeResourceNotFoundException {
			_, err = c.client.CreateLogGroup(&cloudwatchlogs.CreateLogGroupInput{
				LogGroupName: c.logGroupName,
			})
			if err != nil {
				return nil, errors.Wrap(err, "cloudwatchlog.Client.CreateLogGroup")
			}
			return c.getOrCreateLogStream()
		}

		return nil, errors.Wrap(err, "cloudwatchlogs.Client.DescribeLogStreams")
	}

	if len(output.LogStreams) > 0 {
		return output.LogStreams[0], nil
	}

	// No matching log stream, so we need to create it
	_, err = c.client.CreateLogStream(&cloudwatchlogs.CreateLogStreamInput{
		LogGroupName:  c.logGroupName,
		LogStreamName: c.logStreamName,
	})
	if err != nil {
		return nil, errors.Wrap(err, "cloudwatchlogs.Client.CreateLogStream")
	}

	// We can just return an empty log stream as the initial sequence token would be nil anyway.
	return &cloudwatchlogs.LogStream{}, nil
}
