package ai

import (
	"bytes"
	"context"
	"encoding/json"
	"fmt"
	"io"
	"net/http"
	"os"
	"path/filepath"
	"strings"
	"sync"
	"time"

	"pinscrape-allgo/internal/config"
	"pinscrape-allgo/internal/metrics"
)

type KeyPool struct {
	keys  []string
	idx   int
	mu    sync.Mutex
}

func NewKeyPool(keys []string) *KeyPool {
	return &KeyPool{keys: keys}
}

func (p *KeyPool) Next() string {
	p.mu.Lock()
	defer p.mu.Unlock()
	k := p.keys[p.idx%len(p.keys)]
	p.idx++
	return k
}

func (p *KeyPool) Len() int {
	return len(p.keys)
}

func LoadKeys(cfg *config.Config) []string {
	if cfg.AI.KeyFile != "" {
		data, err := os.ReadFile(filepath.Join(cfg.Root, cfg.AI.KeyFile))
		if err == nil {
			var keys []string
			for _, line := range strings.Split(string(data), "\n") {
				line = strings.TrimSpace(line)
				if line != "" && !strings.HasPrefix(line, "#") {
					keys = append(keys, line)
				}
			}
			if len(keys) > 0 {
				return keys
			}
		}
	}
	if cfg.AI.APIKey != "" {
		return []string{cfg.AI.APIKey}
	}
	return nil
}

type ModelInfo struct {
	Model   string
	BaseURL string
	APIKey  string
}

func BuildModelChain(cfg *config.Config) []ModelInfo {
	primary := ModelInfo{
		Model:   cfg.AI.Model,
		BaseURL: cfg.AI.BaseURL,
		APIKey:  cfg.AI.APIKey,
	}
	chain := []ModelInfo{primary}
	for _, fb := range cfg.AI.FallbackModels {
		chain = append(chain, ModelInfo{
			Model:   fb,
			BaseURL: primary.BaseURL,
			APIKey:  primary.APIKey,
		})
	}
	return chain
}

type Client struct {
	cfg     *config.Config
	keys    *KeyPool
	chain   []ModelInfo
	http    *http.Client
	metrics *metrics.Metrics
}

func New(cfg *config.Config, keys []string, chain []ModelInfo, m *metrics.Metrics) *Client {
	transport := &http.Transport{
		MaxIdleConns:        cfg.AI.Concurrency * 2,
		MaxIdleConnsPerHost: cfg.AI.Concurrency * 2,
		IdleConnTimeout:     90 * time.Second,
	}
	return &Client{
		cfg:     cfg,
		keys:    NewKeyPool(keys),
		chain:   chain,
		http:    &http.Client{Timeout: 180 * time.Second, Transport: transport},
		metrics: m,
	}
}

type chatRequest struct {
	Model    string        `json:"model"`
	Messages []chatMessage `json:"messages"`
}

type chatMessage struct {
	Role    string `json:"role"`
	Content string `json:"content"`
}

type chatResponse struct {
	Choices []struct {
		Message struct {
			Content string `json:"content"`
		} `json:"message"`
	} `json:"choices"`
	Error *struct {
		Message string `json:"message"`
	} `json:"error"`
}

// Generate tries each model in the chain with retries, mirroring the Python
// generate_content semantics (429/5xx/timeout are retried, other errors skip
// the model).
func (c *Client) Generate(ctx context.Context, prompt string) (string, error) {
	var lastErr error
	for _, model := range c.chain {
		for attempt := 1; attempt <= c.cfg.AI.MaxRetries; attempt++ {
			apiKey := model.APIKey
			if c.keys.Len() > 0 {
				apiKey = c.keys.Next()
			}
			content, retryable, err := c.call(ctx, model, apiKey, prompt)
			if err == nil {
				return content, nil
			}
			lastErr = err
			if !retryable {
				break
			}
			wait := time.Duration(c.cfg.AI.RetryDelay*attempt) * time.Second
			select {
			case <-ctx.Done():
				return "", ctx.Err()
			case <-time.After(wait):
			}
		}
	}
	return "", fmt.Errorf("all models failed: %w", lastErr)
}

func (c *Client) call(ctx context.Context, model ModelInfo, apiKey, prompt string) (string, bool, error) {
	payload := chatRequest{
		Model: model.Model,
		Messages: []chatMessage{
			{Role: "user", Content: prompt},
		},
	}
	body, _ := json.Marshal(payload)

	url := strings.TrimRight(model.BaseURL, "/") + "/chat/completions"
	req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
	if err != nil {
		return "", false, err
	}
	req.Header.Set("Content-Type", "application/json")
	req.Header.Set("Authorization", "Bearer "+apiKey)

	c.metrics.Request("ai")
	resp, err := c.http.Do(req)
	if err != nil {
		c.metrics.Fail("ai")
		errStr := strings.ToLower(err.Error())
		retryable := strings.Contains(errStr, "timeout") || strings.Contains(errStr, "timed out") ||
			strings.Contains(errStr, "connection reset") || strings.Contains(errStr, "eof")
		return "", retryable, err
	}
	defer resp.Body.Close()
	respBody, _ := io.ReadAll(resp.Body)

	if resp.StatusCode == 429 || resp.StatusCode >= 500 {
		c.metrics.RateLimited("ai")
		if resp.StatusCode >= 500 {
			c.metrics.Fail("ai")
		}
		return "", true, fmt.Errorf("status %d: %s", resp.StatusCode, truncate(string(respBody), 200))
	}
	if resp.StatusCode != 200 {
		c.metrics.Fail("ai")
		return "", false, fmt.Errorf("status %d: %s", resp.StatusCode, truncate(string(respBody), 200))
	}

	var parsed chatResponse
	if err := json.Unmarshal(respBody, &parsed); err != nil {
		c.metrics.Fail("ai")
		return "", false, err
	}
	if len(parsed.Choices) == 0 {
		c.metrics.Fail("ai")
		return "", false, fmt.Errorf("empty choices in response")
	}
	content := strings.TrimSpace(parsed.Choices[0].Message.Content)
	if content == "" {
		c.metrics.Fail("ai")
		return "", true, fmt.Errorf("empty response content")
	}
	c.metrics.OK("ai")
	return content, false, nil
}

func truncate(s string, n int) string {
	if len(s) <= n {
		return s
	}
	return s[:n] + "..."
}

func LoadPrompt(root, mode string) (string, error) {
	data, err := os.ReadFile(filepath.Join(root, "prompts", mode+".txt"))
	if err != nil {
		return "", err
	}
	return string(data), nil
}

func CleanTitle(title string) string {
	for _, r := range []string{`"`, "\u201c", "\u201d", "\u2018", "\u2019"} {
		title = strings.ReplaceAll(title, r, "")
	}
	return strings.TrimSpace(title)
}
