package proxy

import (
	"encoding/json"
	"fmt"
	"net/http"
	"net/http/httptest"
	"net/url"
	"os"
	"path/filepath"
	"strings"
	"sync/atomic"
	"testing"
	"time"

	"github.com/gorilla/websocket"
	"github.com/manaflow-ai/subrouter/internal/accounts"
	"github.com/manaflow-ai/subrouter/internal/authority"
	"github.com/manaflow-ai/subrouter/session"
)

func TestAuthorityRouteStatusScopeAndSelectorConflict(t *testing.T) {
	dir := t.TempDir()
	if err := os.Chmod(dir, 0700); err != nil {
		t.Fatal(err)
	}
	store, err := authority.Open(dir)
	if err != nil {
		t.Fatal(err)
	}
	route := authority.Route{RouteID: authority.NewID(), TenantID: authority.NewID(), Provider: "claude", AccountID: authority.NewID(), Enabled: true}
	wrong := authority.Route{RouteID: authority.NewID(), TenantID: route.TenantID, Provider: "claude", AccountID: authority.NewID(), Enabled: true}
	if err := store.CreateRoute(route); err != nil {
		t.Fatal(err)
	}
	if err := store.CreateRoute(wrong); err != nil {
		t.Fatal(err)
	}
	if err := store.AddFixtureAccount(authority.FixtureAccount{Provider: route.Provider, AccountID: route.AccountID, State: "unavailable"}); err != nil {
		t.Fatal(err)
	}
	key := "synthetic-route-key"
	grant, err := store.CreateGrant(route.RouteID, "workstation", time.Now().Add(time.Hour), authority.HashKey(key))
	if err != nil {
		t.Fatal(err)
	}
	var providerCalls atomic.Int32
	handler := (Server{AuthorityRoutes: store, Transport: authorityRoundTripFunc(func(*http.Request) (*http.Response, error) {
		providerCalls.Add(1)
		panic("synthetic unavailable route reached provider transport")
	})}).Handler()

	status := authorityRequest(t, handler, route.RouteID, "/_subrouter/status", key, "")
	if status.Code != http.StatusOK {
		t.Fatalf("status code = %d", status.Code)
	}
	var payload map[string]any
	if err := json.Unmarshal(status.Body.Bytes(), &payload); err != nil {
		t.Fatal(err)
	}
	body := status.Body.String()
	for _, forbidden := range []string{route.RouteID, route.TenantID, route.AccountID, key, grant.GrantID, grant.KeyHash} {
		if strings.Contains(body, forbidden) {
			t.Fatalf("status exposed forbidden material")
		}
	}
	if payload["provider"] != "claude" || payload["state"] != "migration-required" || payload["account_availability"] != "unavailable" || len(payload) != 7 {
		t.Fatalf("unexpected status: %v", payload)
	}

	conflict := authorityRequest(t, handler, route.RouteID, "/v1/messages", key, "other-account")
	if conflict.Code != http.StatusConflict {
		t.Fatalf("conflict code = %d", conflict.Code)
	}
	unavailable := authorityRequest(t, handler, route.RouteID, "/v1/messages", key, "")
	if unavailable.Code != http.StatusServiceUnavailable {
		t.Fatalf("unavailable code = %d", unavailable.Code)
	}
	if providerCalls.Load() != 0 {
		t.Fatalf("provider transport calls = %d", providerCalls.Load())
	}
	wrongScope := authorityRequest(t, handler, wrong.RouteID, "/_subrouter/status", key, "")
	if wrongScope.Code != http.StatusUnauthorized {
		t.Fatalf("wrong scope code = %d", wrongScope.Code)
	}
	admin := authorityRequest(t, handler, route.RouteID, "/_subrouter/accounts", key, "")
	unauthenticatedAdmin := authorityRequest(t, handler, route.RouteID, "/_subrouter/accounts", "", "")
	if admin.Code != http.StatusNotFound || unauthenticatedAdmin.Code != admin.Code || unauthenticatedAdmin.Body.String() != admin.Body.String() {
		t.Fatalf("routed control responses differ: authenticated=%d %q unauthenticated=%d %q", admin.Code, admin.Body.String(), unauthenticatedAdmin.Code, unauthenticatedAdmin.Body.String())
	}
	if err := store.RevokeGrant(grant.GrantID); err != nil {
		t.Fatal(err)
	}
	revoked := authorityRequest(t, handler, route.RouteID, "/_subrouter/status", key, "")
	if revoked.Code != http.StatusUnauthorized {
		t.Fatalf("revoked code = %d", revoked.Code)
	}
}

func TestAuthorityRouteProxiesFiftyConcurrentWebSockets(t *testing.T) {
	upgrader := websocket.Upgrader{CheckOrigin: func(_ *http.Request) bool { return true }}
	var upstreamHits atomic.Int32
	upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		if got := r.Header.Get("Authorization"); got != "Bearer selected-token" {
			http.Error(w, "wrong provider credential", http.StatusUnauthorized)
			return
		}
		conn, err := upgrader.Upgrade(w, r, nil)
		if err != nil {
			return
		}
		defer conn.Close()
		upstreamHits.Add(1)
		_ = conn.WriteMessage(websocket.TextMessage, []byte("ok"))
	}))
	defer upstream.Close()
	upstreamURL, err := url.Parse(upstream.URL)
	if err != nil {
		t.Fatal(err)
	}

	dir := t.TempDir()
	if err := os.Chmod(dir, 0700); err != nil {
		t.Fatal(err)
	}
	authorityStore, err := authority.Open(dir)
	if err != nil {
		t.Fatal(err)
	}
	defer authorityStore.Close()
	accountID := authority.NewID()
	route := authority.Route{RouteID: authority.NewID(), TenantID: authority.NewID(), Provider: "codex", AccountID: accountID, Enabled: true}
	if err := authorityStore.CreateRoute(route); err != nil {
		t.Fatal(err)
	}
	key := "synthetic-route-key"
	if _, err := authorityStore.CreateGrant(route.RouteID, "workstation", time.Now().Add(24*time.Hour), authority.HashKey(key)); err != nil {
		t.Fatal(err)
	}
	sessions, err := session.NewStore(filepath.Join(t.TempDir(), "sessions.json"))
	if err != nil {
		t.Fatal(err)
	}
	handler := Server{
		AuthorityRoutes: authorityStore,
		Upstream:        upstreamURL,
		Accounts: []accounts.Account{{
			ID:       accountID,
			Provider: accounts.ProviderCodex,
			AuthMode: accounts.AuthModeOAuth,
			Token:    "selected-token",
		}},
		Sessions: sessions,
	}.Handler()
	subrouter := httptest.NewServer(handler)
	defer subrouter.Close()

	const clients = 50
	start := make(chan struct{})
	results := make(chan error, clients)
	wsURL := "ws" + strings.TrimPrefix(subrouter.URL, "http") + "/r/" + route.RouteID + "/v1/responses"
	for i := 0; i < clients; i++ {
		go func(index int) {
			<-start
			header := http.Header{"Authorization": []string{"Bearer " + key}}
			conn, response, err := websocket.DefaultDialer.Dial(wsURL, header)
			if response != nil && response.Body != nil {
				defer response.Body.Close()
			}
			if err != nil {
				results <- fmt.Errorf("client %d handshake: %w", index, err)
				return
			}
			defer conn.Close()
			_, body, err := conn.ReadMessage()
			if err != nil {
				results <- fmt.Errorf("client %d read: %w", index, err)
				return
			}
			if string(body) != "ok" {
				results <- fmt.Errorf("client %d message = %q", index, string(body))
				return
			}
			results <- nil
		}(i)
	}
	close(start)
	for i := 0; i < clients; i++ {
		if err := <-results; err != nil {
			t.Fatal(err)
		}
	}
	if got := upstreamHits.Load(); got != clients {
		t.Fatalf("upstream websocket handshakes = %d, want %d", got, clients)
	}
}


func TestAuthorityAccountIsResolvedExactlyOnceAndBypassesPool(t *testing.T) {
	dir := t.TempDir()
	if err := os.Chmod(dir, 0700); err != nil {
		t.Fatal(err)
	}
	store, err := authority.Open(dir)
	if err != nil {
		t.Fatal(err)
	}
	defer store.Close()
	accountID := authority.NewID()
	route := authority.Route{RouteID: authority.NewID(), TenantID: authority.NewID(), Provider: "claude", AccountID: accountID, Enabled: true}
	request := httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(`{"model":"synthetic"}`))
	exact := accounts.Account{ID: accountID, Provider: accounts.ProviderClaude, AuthMode: accounts.AuthModeOAuth, Token: "synthetic-token"}
	server := Server{AuthorityRoutes: store, Accounts: []accounts.Account{exact}}
	resolved, availability, err := server.resolveAuthorityAccount(request, route)
	if err != nil || availability != "available" || resolved == nil || resolved.ID != exact.ID || resolved.Token != exact.Token {
		t.Fatalf("exact resolution account=%v availability=%q err=%v", resolved, availability, err)
	}
	sessions, err := session.NewStore(filepath.Join(t.TempDir(), "sessions.json"))
	if err != nil {
		t.Fatal(err)
	}
	server.Sessions = sessions
	// Scheduler remains nil: selection must return the immutable authority account
	// before entering any pooled scheduling path.
	selected, _, _, err := server.accountForSessionProviderWithOptions(
		accounts.ProviderClaude, "claude", authority.NewID(),
		request.WithContext(withAuthorityAccount(request.Context(), *resolved)), accountSelectionOptions{},
	)
	if err != nil || selected.ID != exact.ID || selected.Token != exact.Token {
		t.Fatalf("immutable account selection account=%v err=%v", selected, err)
	}

	server.Accounts = []accounts.Account{exact, exact}
	resolved, availability, err = server.resolveAuthorityAccount(request, route)
	if err != nil || resolved != nil || availability != "ambiguous" {
		t.Fatalf("duplicate resolution account=%v availability=%q err=%v", resolved, availability, err)
	}
}

type authorityRoundTripFunc func(*http.Request) (*http.Response, error)

func (fn authorityRoundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) {
	return fn(request)
}

func authorityRequest(t *testing.T, handler http.Handler, routeID, suffix, key, selector string) *httptest.ResponseRecorder {
	t.Helper()
	request := httptest.NewRequest(http.MethodGet, "/r/"+routeID+suffix, nil)
	request.Header.Set("Authorization", "Bearer "+key)
	if selector != "" {
		request.Header.Set("X-Subrouter-Account-ID", selector)
	}
	response := httptest.NewRecorder()
	handler.ServeHTTP(response, request)
	return response
}
