// SPDX-FileCopyrightText: Amolith <amolith@secluded.site>
//
// SPDX-License-Identifier: LicenseRef-MutuaL-1.2

package mcp

import (
	"context"
	"crypto/subtle"
	"errors"
	"fmt"
	"net"
	"net/http"
	"strings"
	"time"

	sdk "github.com/modelcontextprotocol/go-sdk/mcp"
)

const (
	httpReadHeaderTimeout = 10 * time.Second
	httpShutdownTimeout   = 5 * time.Second
)

// RunHTTP serves MCP over Streamable HTTP.
func (s *Server) RunHTTP(ctx context.Context, addr, token string) error {
	listener, err := new(net.ListenConfig).Listen(ctx, "tcp", addr)
	if err != nil {
		return fmt.Errorf("listen MCP HTTP address %q: %w", addr, err)
	}

	return s.serveHTTP(ctx, listener, token)
}

func (s *Server) serveHTTP(ctx context.Context, listener net.Listener, token string) error {
	server := &http.Server{
		Handler:           s.HTTPHandler(token),
		ReadHeaderTimeout: httpReadHeaderTimeout,
	}
	serveErr := make(chan error, 1)
	go func() {
		serveErr <- server.Serve(listener)
	}()

	select {
	case err := <-serveErr:
		if errors.Is(err, http.ErrServerClosed) {
			return nil
		}

		return err
	case <-ctx.Done():
		shutdownCtx, cancel := context.WithTimeout(context.Background(), httpShutdownTimeout)
		defer cancel()

		if err := server.Shutdown(shutdownCtx); err != nil {
			closeErr := server.Close()
			serveError := <-serveErr
			if serveError != nil && errors.Is(serveError, http.ErrServerClosed) {
				serveError = nil
			}

			return errors.Join(ctx.Err(), err, closeErr, serveError)
		}

		if err := <-serveErr; err != nil && !errors.Is(err, http.ErrServerClosed) {
			return errors.Join(ctx.Err(), err)
		}

		return ctx.Err()
	}
}

// HTTPHandler returns the authenticated Streamable HTTP MCP handler.
func (s *Server) HTTPHandler(token string) http.Handler {
	streamable := sdk.NewStreamableHTTPHandler(func(*http.Request) *sdk.Server {
		return s.sdk
	}, nil)

	return newHTTPHandler(token, streamable)
}

func newHTTPHandler(token string, mcpHandler http.Handler) http.Handler {
	mux := http.NewServeMux()
	mux.Handle("/mcp", requireBearerToken(token, mcpHandler))

	return mux
}

func requireBearerToken(token string, next http.Handler) http.Handler {
	if token == "" {
		return next
	}

	return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		scheme, value, ok := strings.Cut(r.Header.Get("Authorization"), " ")
		if !ok || !strings.EqualFold(scheme, "Bearer") ||
			subtle.ConstantTimeCompare([]byte(value), []byte(token)) != 1 {
			w.Header().Set("WWW-Authenticate", "Bearer")
			http.Error(w, "unauthorized", http.StatusUnauthorized)
			return
		}

		next.ServeHTTP(w, r)
	})
}
