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

package mcp

import (
	"context"
	"errors"
	"io"
	"net"
	"net/http"
	"net/http/httptest"
	"strings"
	"testing"
	"time"
)

// server.TRANSPORT.3
func TestRunHTTPACIDServerTransport3ReturnsListenErrorForUnavailableConfiguredAddress(t *testing.T) {
	listener, err := new(net.ListenConfig).Listen(context.Background(), "tcp", "127.0.0.1:0")
	if err != nil {
		t.Fatalf("listen occupied address: %v", err)
	}
	t.Cleanup(func() {
		if closeErr := listener.Close(); closeErr != nil {
			t.Fatalf("close listener: %v", closeErr)
		}
	})

	err = NewServer(nil, "test").RunHTTP(context.Background(), listener.Addr().String(), "")
	if err == nil {
		t.Fatal("RunHTTP() error = nil, want listen error")
	}
	if !strings.Contains(err.Error(), "listen MCP HTTP address") {
		t.Fatalf("RunHTTP() error = %v, want listen address error", err)
	}
}

// server.TRANSPORT.3
func TestServeHTTPACIDServerTransport3ServesConfiguredListenerUntilContextCanceled(t *testing.T) {
	listener, err := new(net.ListenConfig).Listen(context.Background(), "tcp", "127.0.0.1:0")
	if err != nil {
		t.Fatalf("listen configured address: %v", err)
	}

	ctx, cancel := context.WithCancel(context.Background())
	serveErr := make(chan error, 1)
	go func() {
		serveErr <- NewServer(nil, "test").serveHTTP(ctx, listener, "")
	}()

	requestCtx, stopRequest := context.WithTimeout(context.Background(), time.Second)
	defer stopRequest()
	request, err := http.NewRequestWithContext(
		requestCtx,
		http.MethodGet,
		"http://"+listener.Addr().String()+"/not-mcp",
		nil,
	)
	if err != nil {
		t.Fatalf("new request: %v", err)
	}

	response, err := http.DefaultClient.Do(request)
	if err != nil {
		cancel()
		t.Fatalf("request configured listener: %v", err)
	}
	if _, err := io.Copy(io.Discard, response.Body); err != nil {
		cancel()
		t.Fatalf("read response body: %v", err)
	}
	if err := response.Body.Close(); err != nil {
		cancel()
		t.Fatalf("close response body: %v", err)
	}
	if response.StatusCode != http.StatusNotFound {
		cancel()
		t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusNotFound)
	}

	cancel()
	select {
	case err := <-serveErr:
		if !errors.Is(err, context.Canceled) {
			t.Fatalf("serveHTTP() error = %v, want context.Canceled", err)
		}
	case <-time.After(time.Second):
		t.Fatal("serveHTTP() did not stop after context cancellation")
	}
}

// server.TRANSPORT.5
func TestHTTPHandlerACIDServerTransport5RoutesMCPPath(t *testing.T) {
	calls := 0
	handler := newHTTPHandler("", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
		calls++
		w.WriteHeader(http.StatusNoContent)
	}))

	recorder := httptest.NewRecorder()
	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
	handler.ServeHTTP(recorder, request)

	if recorder.Code != http.StatusNoContent {
		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNoContent)
	}
	if calls != 1 {
		t.Fatalf("MCP handler calls = %d, want 1", calls)
	}
}

// server.TRANSPORT.5
func TestHTTPHandlerACIDServerTransport5DoesNotRouteOtherPaths(t *testing.T) {
	tests := []string{"/", "/mcp/extra"}

	for _, path := range tests {
		t.Run(path, func(t *testing.T) {
			calls := 0
			handler := newHTTPHandler("", http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
				calls++
			}))

			recorder := httptest.NewRecorder()
			request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, path, nil)
			handler.ServeHTTP(recorder, request)

			if recorder.Code != http.StatusNotFound {
				t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNotFound)
			}
			if calls != 0 {
				t.Fatalf("MCP handler calls = %d, want 0", calls)
			}
		})
	}
}

// server.SECURITY.7
func TestHTTPHandlerACIDServerSecurity7GatesMCPPathWithBearerToken(t *testing.T) {
	calls := 0
	handler := newHTTPHandler("secret-token", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
		calls++
		w.WriteHeader(http.StatusNoContent)
	}))

	unauthorized := httptest.NewRecorder()
	unauthorizedRequest := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
	handler.ServeHTTP(unauthorized, unauthorizedRequest)

	if unauthorized.Code != http.StatusUnauthorized {
		t.Fatalf("unauthorized status = %d, want %d", unauthorized.Code, http.StatusUnauthorized)
	}
	if calls != 0 {
		t.Fatalf("MCP handler calls after unauthorized request = %d, want 0", calls)
	}

	authorized := httptest.NewRecorder()
	authorizedRequest := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
	authorizedRequest.Header.Set("Authorization", "Bearer secret-token")
	handler.ServeHTTP(authorized, authorizedRequest)

	if authorized.Code != http.StatusNoContent {
		t.Fatalf("authorized status = %d, want %d", authorized.Code, http.StatusNoContent)
	}
	if calls != 1 {
		t.Fatalf("MCP handler calls after authorized request = %d, want 1", calls)
	}
}

// server.SECURITY.7
func TestServerHTTPHandlerACIDServerSecurity7WrapsSDKHandlerWithBearerAuth(t *testing.T) {
	handler := NewServer(nil, "test").HTTPHandler("secret-token")

	recorder := httptest.NewRecorder()
	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
	handler.ServeHTTP(recorder, request)

	if recorder.Code != http.StatusUnauthorized {
		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusUnauthorized)
	}
}

// server.SECURITY.3 server.SECURITY.6
func TestHTTPBearerAuthACIDServerSecurity3And6AllowsMatchingBearerToken(t *testing.T) {
	calls := 0
	handler := requireBearerToken("secret-token", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
		calls++
		w.WriteHeader(http.StatusNoContent)
	}))

	recorder := httptest.NewRecorder()
	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
	request.Header.Set("Authorization", "Bearer secret-token")
	handler.ServeHTTP(recorder, request)

	if recorder.Code != http.StatusNoContent {
		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNoContent)
	}
	if calls != 1 {
		t.Fatalf("next calls = %d, want 1", calls)
	}
}

// server.SECURITY.7
func TestHTTPBearerAuthACIDServerSecurity7RejectsInvalidTokens(t *testing.T) {
	tests := []struct {
		name   string
		header string
	}{
		{name: "missing", header: ""},
		{name: "wrong", header: "Bearer wrong-token"},
		{name: "basic", header: "Basic secret-token"},
		{name: "malformed", header: "Bearer"},
	}

	for _, tt := range tests {
		t.Run(tt.name, func(t *testing.T) {
			calls := 0
			handler := requireBearerToken("secret-token", http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
				calls++
			}))

			recorder := httptest.NewRecorder()
			request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
			if tt.header != "" {
				request.Header.Set("Authorization", tt.header)
			}
			handler.ServeHTTP(recorder, request)

			if recorder.Code != http.StatusUnauthorized {
				t.Fatalf("status = %d, want %d", recorder.Code, http.StatusUnauthorized)
			}
			if got := recorder.Header().Get("WWW-Authenticate"); got != "Bearer" {
				t.Fatalf("WWW-Authenticate = %q, want Bearer", got)
			}
			if calls != 0 {
				t.Fatalf("next calls = %d, want 0", calls)
			}
		})
	}
}

func TestHTTPBearerAuthAllowsRequestsWhenNoTokenConfigured(t *testing.T) {
	calls := 0
	handler := requireBearerToken("", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
		calls++
		w.WriteHeader(http.StatusNoContent)
	}))

	recorder := httptest.NewRecorder()
	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
	handler.ServeHTTP(recorder, request)

	if recorder.Code != http.StatusNoContent {
		t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNoContent)
	}
	if calls != 1 {
		t.Fatalf("next calls = %d, want 1", calls)
	}
}

// server.SECURITY.2
func TestHTTPBearerAuthACIDServerSecurity2DoesNotExposeTokensInUnauthorizedResponse(t *testing.T) {
	handler := requireBearerToken("secret-token", http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
		t.Fatal("next handler was called")
	}))

	recorder := httptest.NewRecorder()
	request := httptest.NewRequestWithContext(context.Background(), http.MethodPost, "/mcp", nil)
	request.Header.Set("Authorization", "Bearer submitted-token")
	handler.ServeHTTP(recorder, request)

	body := recorder.Body.String()
	if strings.Contains(body, "secret-token") {
		t.Fatalf("response body exposed configured token: %q", body)
	}
	if strings.Contains(body, "submitted-token") {
		t.Fatalf("response body exposed submitted token: %q", body)
	}
}
