// Package templates provides NetFlow/IPFIX template system helpers.
package templates

import (
	"fmt"
	"sync"
	"time"

	"github.com/netsampler/goflow2/v3/decoders/netflow"
)

// TemplateKey identifies a template entry for expiry tracking.
type TemplateKey struct {
	Version     uint16
	ObsDomainID uint32
	TemplateID  uint16
}

// ExpiringRegistry provides expiry controls for all router template systems.
type ExpiringRegistry struct {
	lock           sync.RWMutex                       // protects registry state and sweeper lifecycle fields
	wrapped        Registry                           // underlying registry to store templates
	systems        map[string]*ExpiringTemplateSystem // per-router template systems
	counts         map[string]int                     // template counts per router
	emptySince     map[string]time.Time               // timestamp when a router became empty
	ttl            time.Duration                      // template TTL for expiry
	now            func() time.Time                   // time source (overridable for tests)
	extendOnAccess bool                               // refresh expiry on GetTemplate when true
	sweepInterval  time.Duration                      // sweep interval used on Start
	sweeperStop    chan struct{}                      // signals sweeper goroutine to stop
	sweeperDone    chan struct{}                      // closed when sweeper goroutine exits
	closeOnce      sync.Once                          // guards Close
	startOnce      sync.Once                          // guards Start
}

// ExpiringOption configures an ExpiringRegistry.
type ExpiringOption func(*ExpiringRegistry)

// WithExtendOnAccess toggles whether template accesses refresh their expiry time.
func WithExtendOnAccess(enable bool) ExpiringOption {
	return func(r *ExpiringRegistry) {
		r.extendOnAccess = enable
	}
}

// WithSweepInterval configures the sweeper interval used on Start.
func WithSweepInterval(interval time.Duration) ExpiringOption {
	return func(r *ExpiringRegistry) {
		r.sweepInterval = interval
	}
}

// NewExpiringRegistry wraps a registry with expiry tracking.
func NewExpiringRegistry(wrapped Registry, ttl time.Duration, opts ...ExpiringOption) *ExpiringRegistry {
	if wrapped == nil {
		wrapped = NewInMemoryRegistry(nil)
	}
	registry := &ExpiringRegistry{
		wrapped:    wrapped,
		systems:    make(map[string]*ExpiringTemplateSystem),
		counts:     make(map[string]int),
		emptySince: make(map[string]time.Time),
		ttl:        ttl,
		now:        time.Now,
	}
	for _, opt := range opts {
		if opt != nil {
			opt(registry)
		}
	}
	return registry
}

// GetSystem returns a wrapped template system for a router.
func (r *ExpiringRegistry) GetSystem(key string) netflow.NetFlowTemplateSystem {
	r.lock.RLock()
	system, ok := r.systems[key]
	r.lock.RUnlock()
	if ok {
		r.lock.Lock()
		if r.counts[key] == 0 {
			// Empty-system tracking does not follow extendOnAccess and refreshes on lookup,
			// which helps avoid churn from repeated lookups of unknown packets.
			r.emptySince[key] = r.now()
		}
		r.lock.Unlock()
		return system
	}

	wrapped := r.wrapped.GetSystem(key)
	system = NewExpiringTemplateSystem(wrapped, r.ttl)
	system.now = r.now
	system.key = key
	system.reg = r
	system.extendOnAccess = r.extendOnAccess

	r.lock.Lock()
	if existing, ok := r.systems[key]; ok {
		r.lock.Unlock()
		return existing
	}
	r.systems[key] = system
	count := len(wrapped.GetTemplates())
	r.counts[key] = count
	if count == 0 {
		r.emptySince[key] = r.now()
	} else {
		delete(r.emptySince, key)
	}
	r.lock.Unlock()
	return system
}

// GetAll returns all templates for every router.
func (r *ExpiringRegistry) GetAll() map[string]netflow.FlowBaseTemplateSet {
	r.lock.RLock()
	systems := make(map[string]netflow.NetFlowTemplateSystem, len(r.systems))
	for key, system := range r.systems {
		systems[key] = system
	}
	r.lock.RUnlock()

	ret := make(map[string]netflow.FlowBaseTemplateSet, len(systems))
	for key, system := range systems {
		ret[key] = system.GetTemplates()
	}
	return ret
}

func (r *ExpiringRegistry) increment(key string) {
	r.lock.Lock()
	count := r.counts[key] + 1
	r.counts[key] = count
	if count == 1 {
		delete(r.emptySince, key)
	}
	r.lock.Unlock()
}

func (r *ExpiringRegistry) decrement(key string) int {
	r.lock.Lock()
	count := r.counts[key] - 1
	if count < 0 {
		count = 0
	}
	r.counts[key] = count
	if count == 0 {
		r.emptySince[key] = r.now()
	}
	r.lock.Unlock()
	return count
}

func (r *ExpiringRegistry) pruneEmptyBefore(cutoff time.Time) int {
	pruned := 0
	r.lock.Lock()
	for key, since := range r.emptySince {
		if since.Before(cutoff) {
			delete(r.emptySince, key)
			delete(r.counts, key)
			delete(r.systems, key)
			r.wrapped.RemoveSystem(key)
			pruned++
		}
	}
	r.lock.Unlock()
	return pruned
}

// ExpireBefore removes templates older than the cutoff across all routers.
func (r *ExpiringRegistry) ExpireBefore(cutoff time.Time) int {
	r.lock.RLock()
	systems := make([]struct {
		key    string
		system *ExpiringTemplateSystem
	}, 0, len(r.systems))
	for key, system := range r.systems {
		systems = append(systems, struct {
			key    string
			system *ExpiringTemplateSystem
		}{
			key:    key,
			system: system,
		})
	}
	r.lock.RUnlock()

	removed := 0
	for _, entry := range systems {
		removed += entry.system.ExpireBefore(cutoff)
	}
	return removed
}

// ExpireStale removes stale templates and prunes empty systems using the TTL.
func (r *ExpiringRegistry) ExpireStale() (int, int) {
	removed := 0
	if r.ttl > 0 {
		removed = r.ExpireBefore(r.now().Add(-r.ttl))
	}
	// When template expiry is disabled, still prune empty systems on sweeps.
	pruned := r.pruneEmptyBefore(r.now().Add(-r.ttl))
	return removed, pruned
}

// StartSweeper begins periodic expiry using the configured TTL.
// StartSweeper starts the periodic expiry/empty cleanup loop once; repeated calls are no-ops.
func (r *ExpiringRegistry) StartSweeper(interval time.Duration) {
	if interval <= 0 {
		return
	}

	r.lock.Lock()
	r.sweepInterval = interval
	// Only one sweeper should be active; return if already started.
	if r.sweeperStop != nil {
		r.lock.Unlock()
		return
	}
	r.sweeperStop = make(chan struct{})
	r.sweeperDone = make(chan struct{})
	stop := r.sweeperStop
	done := r.sweeperDone
	r.lock.Unlock()

	go func() {
		ticker := time.NewTicker(interval)
		defer ticker.Stop()
		defer close(done)
		for {
			select {
			case <-ticker.C:
				r.ExpireStale()
			case <-stop:
				return
			}
		}
	}()
}

// Close stops the sweeper goroutine if it is running.
func (r *ExpiringRegistry) Close() {
	r.closeOnce.Do(func() {
		r.lock.Lock()
		if r.sweeperStop == nil {
			r.lock.Unlock()
			r.wrapped.Close()
			return
		}
		stop := r.sweeperStop
		done := r.sweeperDone
		r.sweeperStop = nil
		r.sweeperDone = nil
		r.lock.Unlock()

		close(stop)
		<-done
		r.wrapped.Close()
	})
}

// Start forwards start to the wrapped registry.
// Start initializes the wrapped registry and sweeper once; repeated calls are no-ops.
func (r *ExpiringRegistry) Start() {
	r.startOnce.Do(func() {
		r.lock.Lock()
		interval := r.sweepInterval
		r.lock.Unlock()
		r.wrapped.Start()
		if interval > 0 {
			r.StartSweeper(interval)
		}
	})
}

// RemoveSystem deletes a router entry if present.
func (r *ExpiringRegistry) RemoveSystem(key string) {
	r.lock.Lock()
	delete(r.systems, key)
	delete(r.counts, key)
	delete(r.emptySince, key)
	r.wrapped.RemoveSystem(key)
	r.lock.Unlock()
}

// ExpiringTemplateSystem tracks template update times and supports expiry.
type ExpiringTemplateSystem struct {
	wrapped        netflow.NetFlowTemplateSystem
	lock           sync.Mutex
	updated        map[TemplateKey]time.Time
	ttl            time.Duration
	now            func() time.Time
	key            string
	reg            *ExpiringRegistry
	extendOnAccess bool
}

// NewExpiringTemplateSystem wraps a template system with expiry tracking.
func NewExpiringTemplateSystem(wrapped netflow.NetFlowTemplateSystem, ttl time.Duration) *ExpiringTemplateSystem {
	if wrapped == nil {
		wrapped = netflow.CreateTemplateSystem()
	}
	return &ExpiringTemplateSystem{
		wrapped: wrapped,
		updated: make(map[TemplateKey]time.Time),
		ttl:     ttl,
		now:     time.Now,
	}
}

// AddTemplate records template update time and forwards to the wrapped system.
func (s *ExpiringTemplateSystem) AddTemplate(version uint16, obsDomainId uint32, templateId uint16, template interface{}) (netflow.TemplateStatus, error) {
	s.lock.Lock()
	update, err := s.wrapped.AddTemplate(version, obsDomainId, templateId, template)
	if err != nil {
		s.lock.Unlock()
		return update, fmt.Errorf("expiring templates add %d/%d/%d: %w", version, obsDomainId, templateId, err)
	}
	s.updated[TemplateKey{Version: version, ObsDomainID: obsDomainId, TemplateID: templateId}] = s.now()
	if s.reg != nil {
		if update == netflow.TemplateAdded {
			s.reg.increment(s.key)
		}
	}
	s.lock.Unlock()
	return update, nil
}

// GetTemplate forwards template lookup to the wrapped system.
func (s *ExpiringTemplateSystem) GetTemplate(version uint16, obsDomainId uint32, templateId uint16) (interface{}, error) {
	s.lock.Lock()
	template, err := s.wrapped.GetTemplate(version, obsDomainId, templateId)
	if err != nil {
		s.lock.Unlock()
		return template, fmt.Errorf("expiring templates get %d/%d/%d: %w", version, obsDomainId, templateId, err)
	}
	if s.extendOnAccess {
		s.updated[TemplateKey{Version: version, ObsDomainID: obsDomainId, TemplateID: templateId}] = s.now()
	}
	s.lock.Unlock()
	return template, nil
}

// RemoveTemplate removes a template and its tracking entry.
func (s *ExpiringTemplateSystem) RemoveTemplate(version uint16, obsDomainId uint32, templateId uint16) (interface{}, bool, error) {
	s.lock.Lock()
	template, removed, err := s.wrapped.RemoveTemplate(version, obsDomainId, templateId)
	if removed {
		delete(s.updated, TemplateKey{Version: version, ObsDomainID: obsDomainId, TemplateID: templateId})
		if s.reg != nil {
			s.reg.decrement(s.key)
		}
	}
	s.lock.Unlock()
	if err != nil {
		return template, removed, fmt.Errorf("expiring templates remove %d/%d/%d: %w", version, obsDomainId, templateId, err)
	}
	return template, removed, nil
}

// GetTemplates returns all templates from the wrapped system.
func (s *ExpiringTemplateSystem) GetTemplates() netflow.FlowBaseTemplateSet {
	return s.wrapped.GetTemplates()
}

// ExpireBefore removes templates last updated before the cutoff.
func (s *ExpiringTemplateSystem) ExpireBefore(cutoff time.Time) int {
	removed := 0
	s.lock.Lock()
	for key, updated := range s.updated {
		if !updated.Before(cutoff) {
			continue
		}
		if _, removedTemplate, err := s.wrapped.RemoveTemplate(key.Version, key.ObsDomainID, key.TemplateID); err == nil && removedTemplate {
			delete(s.updated, key)
			if s.reg != nil {
				s.reg.decrement(s.key)
			}
			removed++
		}
	}
	s.lock.Unlock()
	return removed
}

// ExpireStale removes templates older than the configured TTL.
func (s *ExpiringTemplateSystem) ExpireStale() int {
	if s.ttl <= 0 {
		return 0
	}
	return s.ExpireBefore(s.now().Add(-s.ttl))
}
