package services

import (
	"context"
	"errors"
	"fmt"
	"strings"
	"sync"
	"sync/atomic"
	"testing"
	"time"

	"github.com/cloudwego/eino/components/model"
	"github.com/cloudwego/eino/schema"
	"wa-assistant/backend/models"
)

// These fixtures never open a network connection or execute a business tool.
type runtimeCountingModel struct {
	calls *atomic.Int32
	after func()
}

func (m *runtimeCountingModel) WithTools([]*schema.ToolInfo) (model.ToolCallingChatModel, error) {
	return &runtimeCountingModel{calls: m.calls, after: m.after}, nil
}
func (m *runtimeCountingModel) Generate(context.Context, []*schema.Message, ...model.Option) (*schema.Message, error) {
	m.calls.Add(1)
	if m.after != nil {
		m.after()
	}
	return schema.AssistantMessage(`{"supported":true,"unsupported_claims":[],"conflict":false,"complete":true,"missing_aspects":[]}`, nil), nil
}
func (m *runtimeCountingModel) Stream(context.Context, []*schema.Message, ...model.Option) (*schema.StreamReader[*schema.Message], error) {
	return nil, errors.New("fixture stream unused")
}

func TestAgentRuntimeRejectsCanceledModelCallBeforeProvider(t *testing.T) {
	var calls atomic.Int32
	m := &observedChatModel{inner: &runtimeCountingModel{calls: &calls}}
	ctx, cancel := context.WithCancel(context.Background())
	cancel()
	out, err := m.Generate(ctx, []*schema.Message{schema.UserMessage("Halo")})
	if !errors.Is(err, context.Canceled) || out != nil || calls.Load() != 0 {
		t.Fatalf("canceled request reached provider: calls=%d result=%v err=%v", calls.Load(), out, err)
	}
}

func TestAgentRuntimeDiscardsModelResultAfterCancellation(t *testing.T) {
	var calls atomic.Int32
	ctx, cancel := context.WithCancel(context.Background())
	defer cancel()
	m := &observedChatModel{inner: &runtimeCountingModel{calls: &calls, after: cancel}}
	out, err := m.Generate(ctx, []*schema.Message{schema.UserMessage("Halo")})
	if !errors.Is(err, context.Canceled) || out != nil || calls.Load() != 1 {
		t.Fatalf("late result survived cancellation: calls=%d result=%v err=%v", calls.Load(), out, err)
	}
}

func TestAgentRuntimeBudgetSharedByBoundToolsAndReviewer(t *testing.T) {
	var calls atomic.Int32
	m := &observedChatModel{inner: &runtimeCountingModel{calls: &calls}, budget: &agentModelBudget{}}
	bound, err := m.WithTools([]*schema.ToolInfo{{Name: "search_knowledge"}})
	if err != nil {
		t.Fatal(err)
	}
	for i := 0; i < 13; i++ {
		if _, err := bound.Generate(context.Background(), []*schema.Message{schema.UserMessage("FAQ")}); err != nil {
			t.Fatalf("legitimate call %d rejected: %v", i+1, err)
		}
	}
	review := modelAnswerReviewer(m)
	for i := 0; i < 3; i++ {
		if _, err := review(context.Background(), "FAQ", "Fakta", "Fakta"); err != nil {
			t.Fatalf("review %d rejected: %v", i+1, err)
		}
	}
	out, err := bound.Generate(context.Background(), []*schema.Message{schema.UserMessage("FAQ")})
	if err == nil || !strings.Contains(err.Error(), "model_call_budget_exhausted") || out != nil || calls.Load() != 16 {
		t.Fatalf("nested call bypassed turn budget: calls=%d result=%v err=%v", calls.Load(), out, err)
	}
}

func TestAgentRuntimeBudgetIsAtomicForConcurrentCalls(t *testing.T) {
	var calls atomic.Int32
	m := &observedChatModel{inner: &runtimeCountingModel{calls: &calls}, budget: &agentModelBudget{}}
	var wg sync.WaitGroup
	for i := 0; i < 32; i++ {
		wg.Add(1)
		go func() {
			defer wg.Done()
			_, _ = m.Generate(context.Background(), []*schema.Message{schema.UserMessage("FAQ")})
		}()
	}
	wg.Wait()
	if got := calls.Load(); got != 16 {
		t.Fatalf("concurrent calls bypassed budget: %d", got)
	}
}

func TestAgentRuntimeBudgetIsolatedAcrossTurns(t *testing.T) {
	for turn := 0; turn < 2; turn++ {
		var calls atomic.Int32
		m := &observedChatModel{inner: &runtimeCountingModel{calls: &calls}, budget: &agentModelBudget{}}
		for i := 0; i < 17; i++ {
			_, _ = m.Generate(context.Background(), []*schema.Message{schema.UserMessage("FAQ")})
		}
		if got := calls.Load(); got != 16 {
			t.Fatalf("turn %d has wrong independent budget: calls=%d", turn, got)
		}
	}
	// Knowledge ingestion reuses observedChatModel across independent chunks;
	// its own graph/deadline limits must not inherit a customer-turn allowance.
	var trainingCalls atomic.Int32
	training := &observedChatModel{inner: &runtimeCountingModel{calls: &trainingCalls}}
	for i := 0; i < 20; i++ {
		if _, err := training.Generate(context.Background(), []*schema.Message{schema.UserMessage("training chunk")}); err != nil {
			t.Fatalf("non-chat workflow inherited turn allowance: %v", err)
		}
	}
	if trainingCalls.Load() != 20 {
		t.Fatal("non-chat workflow calls were capped")
	}
}

func TestAgentRuntimeChatPreservesCancellation(t *testing.T) {
	agenticTestDB(t)
	t.Setenv("AI_PROVIDER", "openrouter")
	t.Setenv("OPENROUTER_API_KEY", "fixture-unused-key")
	ctx, cancel := context.WithCancel(context.Background())
	cancel()
	result, err := chatAgentic(models.Agent{ID: 1}, "", "", "Halo", nil, ChatOptions{Context: ctx, DryRun: true})
	if !errors.Is(err, context.Canceled) || result.Reply != "" {
		t.Fatalf("chat erased cancellation: reply=%q err=%v", result.Reply, err)
	}
}

func TestAgentRuntimeChatPreservesExpiredDeadline(t *testing.T) {
	agenticTestDB(t)
	t.Setenv("AI_PROVIDER", "openrouter")
	t.Setenv("OPENROUTER_API_KEY", "fixture-unused-key")
	ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-time.Second))
	defer cancel()
	result, err := chatAgentic(models.Agent{ID: 1}, "", "", "Halo", nil, ChatOptions{Context: ctx, DryRun: true})
	if !errors.Is(err, context.DeadlineExceeded) || result.Reply != "" {
		t.Fatalf("chat erased deadline: reply=%q err=%v", result.Reply, err)
	}
}

func TestAgentRuntimeSafeErrorsKeepClassificationWithoutProviderDetails(t *testing.T) {
	const privateDetail = "fixture-private-provider-detail"
	for _, cause := range []error{context.Canceled, context.DeadlineExceeded, errAgentModelBudget} {
		err := safeAgentRuntimeError(fmt.Errorf("%s: %w", privateDetail, cause))
		if !errors.Is(err, cause) || strings.Contains(err.Error(), privateDetail) {
			t.Fatalf("safe error lost cause or leaked provider detail: %v", err)
		}
	}
	if got := safeAgentRuntimeError(errors.New(privateDetail)).Error(); strings.Contains(got, privateDetail) {
		t.Fatalf("generic provider details leaked: %s", got)
	}
	if got := simulationErrorCode(errAgentModelBudget); got != "model_call_budget_exhausted" {
		t.Fatalf("budget stop misclassified as service error: %s", got)
	}
}
