package quicreuse

import (
	
	
	
	
	
	
	
	

	
	
	ma 
	
)

type Listener interface {
	Accept(context.Context) (*quic.Conn, error)
	Addr() net.Addr
	Multiaddrs() []ma.Multiaddr
	io.Closer
}

type protoConf struct {
	ln                  *listener
	tlsConf             *tls.Config
	allowWindowIncrease func(conn *quic.Conn, delta uint64) bool
}

type quicListener struct {
	l         QUICListener
	transport RefCountedQUICTransport
	running   chan struct{}
	addrs     []ma.Multiaddr

	protocolsMu sync.Mutex
	protocols   map[string]protoConf
}

func newQuicListener( RefCountedQUICTransport,  *quic.Config) (*quicListener, error) {
	 := make([]ma.Multiaddr, 0, 2)
	,  := ToQuicMultiaddr(.LocalAddr(), quic.Version1)
	if  != nil {
		return nil, 
	}
	 = append(, )
	 := &quicListener{
		protocols: map[string]protoConf{},
		running:   make(chan struct{}),
		transport: ,
		addrs:     ,
	}
	 := &tls.Config{
		SessionTicketsDisabled: true, // This is set for the config for client, but we set it here as well: https://github.com/quic-go/quic-go/issues/4029
		GetConfigForClient: func( *tls.ClientHelloInfo) (*tls.Config, error) {
			.protocolsMu.Lock()
			defer .protocolsMu.Unlock()
			for ,  := range .SupportedProtos {
				if ,  := .protocols[];  {
					 := .tlsConf
					if .GetConfigForClient != nil {
						return .GetConfigForClient()
					}
					return , nil
				}
			}
			return nil, fmt.Errorf("no supported protocol found. offered: %+v", .SupportedProtos)
		},
	}
	 := .Clone()
	.AllowConnectionWindowIncrease = .allowWindowIncrease
	,  := .Listen(, )
	if  != nil {
		return nil, 
	}
	.l = 
	go .Run() // This go routine shuts down once the underlying quic.Listener is closed (or returns an error).
	return , nil
}

func ( *quicListener) ( *quic.Conn,  uint64) bool {
	.protocolsMu.Lock()
	defer .protocolsMu.Unlock()

	,  := .protocols[.ConnectionState().TLS.NegotiatedProtocol]
	if ! {
		return false
	}
	return .allowWindowIncrease(, )
}

func ( *quicListener) ( any,  *tls.Config,  func( *quic.Conn,  uint64) bool,  func()) (*listener, error) {
	.protocolsMu.Lock()
	defer .protocolsMu.Unlock()

	if len(.NextProtos) == 0 {
		return nil, errors.New("no ALPN found in tls.Config")
	}

	for ,  := range .NextProtos {
		if ,  := .protocols[];  {
			return nil, fmt.Errorf("already listening for protocol %s", )
		}
	}

	 := &listener{
		queue:             make(chan *quic.Conn, queueLen),
		acceptLoopRunning: .running,
		addr:              .l.Addr(),
		addrs:             .addrs,
	}
	if  != nil {
		if ,  := .transport.(*refcountedTransport);  {
			.associateForListener(, )
		}
	}

	.remove = func() {
		if  != nil {
			if ,  := .transport.(*refcountedTransport);  {
				.RemoveAssociationsForListener()
			}
		}
		.protocolsMu.Lock()
		for ,  := range .NextProtos {
			delete(.protocols, )
		}
		.protocolsMu.Unlock()
		()
	}

	for ,  := range .NextProtos {
		.protocols[] = protoConf{
			ln:                  ,
			tlsConf:             ,
			allowWindowIncrease: ,
		}
	}
	return , nil
}

func ( *quicListener) () error {
	defer close(.running)
	defer .transport.DecreaseCount()
	for {
		,  := .l.Accept(context.Background())
		if  != nil {
			if errors.Is(, quic.ErrServerClosed) || strings.Contains(.Error(), "use of closed network connection") {
				return transport.ErrListenerClosed
			}
			return 
		}
		 := .ConnectionState().TLS.NegotiatedProtocol

		.protocolsMu.Lock()
		,  := .protocols[]
		if ! {
			.protocolsMu.Unlock()
			return fmt.Errorf("negotiated unknown protocol: %s", )
		}
		.ln.add()
		.protocolsMu.Unlock()
	}
}

func ( *quicListener) () error {
	 := .l.Close()
	<-.running // wait for Run to return
	return 
}

const queueLen = 16

// A listener for a single ALPN protocol (set).
type listener struct {
	queue             chan *quic.Conn
	acceptLoopRunning chan struct{}
	addr              net.Addr
	addrs             []ma.Multiaddr
	remove            func()
	closeOnce         sync.Once
}

var _ Listener = &listener{}

func ( *listener) ( *quic.Conn) {
	select {
	case .queue <- :
	default:
		.CloseWithError(1, "queue full")
	}
}

func ( *listener) ( context.Context) (*quic.Conn, error) {
	select {
	case <-.Done():
		return nil, .Err()
	case <-.acceptLoopRunning:
		return nil, transport.ErrListenerClosed
	case ,  := <-.queue:
		if ! {
			return nil, transport.ErrListenerClosed
		}
		return , nil
	}
}

func ( *listener) () net.Addr {
	return .addr
}

func ( *listener) () []ma.Multiaddr {
	return .addrs
}

func ( *listener) () error {
	.closeOnce.Do(func() {
		.remove()
		close(.queue)
		// drain the queue
		for  := range .queue {
			.CloseWithError(quic.ApplicationErrorCode(network.ConnShutdown), "closing")
		}
	})
	return nil
}