package pipeline

import (
	"context"
	"log"
	"path/filepath"
	"strings"
	"sync"
	"sync/atomic"

	"pinscrape-allgo/internal/clients/ai"
	"pinscrape-allgo/internal/db"
	"pinscrape-allgo/internal/util"
)

type aiJob struct {
	row    db.WorkRow
	mode   string
	writer *db.Writer
}

// ScrapeImagesAndQueueAI runs the scraper and the AI generator concurrently:
// every keyword whose images are saved is immediately queued for ai_title and
// ai_content while the scraper continues with the next keywords. Wall time
// approaches max(scrape, ai) instead of scrape + ai. The search is merged:
// descriptions are stored together with the images from the same response.
func ScrapeImagesAndQueueAI(ctx context.Context, env *Env, aiClient *ai.Client) error {
	titlePrompt, err := ai.LoadPrompt(env.Cfg.Root, "title")
	if err != nil {
		return err
	}
	articlePrompt, err := ai.LoadPrompt(env.Cfg.Root, "article")
	if err != nil {
		return err
	}

	files, err := db.ListDBFiles(env.Cfg.Path(env.Cfg.DataFolder))
	if err != nil {
		return err
	}

	var totalAICount int64
	for _, f := range files {
		if ctx.Err() != nil {
			break
		}
		n, err := scrapeAndQueueOneDB(ctx, env, f, aiClient, titlePrompt, articlePrompt)
		if err != nil {
			return err
		}
		totalAICount += n
	}
	log.Printf("AI overlapped phase done: %d items generated", totalAICount)
	return nil
}

func scrapeAndQueueOneDB(ctx context.Context, env *Env, dbPath string, aiClient *ai.Client, titlePrompt, articlePrompt string) (int64, error) {
	conn, err := db.Open(dbPath)
	if err != nil {
		return 0, err
	}
	pending, err := db.GetPendingImagesDesc(conn)
	if err != nil {
		db.Close(conn)
		return 0, err
	}
	rows := make([]db.WorkRow, 0, len(pending))
	for _, r := range pending {
		if r.NeedImages {
			rows = append(rows, r)
		}
	}
	if len(rows) == 0 {
		log.Printf("No pending keywords in %s", filepath.Base(dbPath))
		db.Close(conn)
		return 0, nil
	}
	log.Printf("Found %d pending keywords in %s", len(rows), filepath.Base(dbPath))

	writer := db.NewWriter(conn, db.WriterOptions{
		BatchSize: env.Cfg.Go.WriterBatchSize,
	})

	jobs := make(chan aiJob, env.Cfg.AI.Concurrency*2)
	var aiWG sync.WaitGroup
	var aiCounter int64
	aiSem := make(chan struct{}, env.Cfg.AI.Concurrency)

	var consumers sync.WaitGroup
	consumers.Add(1)
	go func() {
		defer consumers.Done()
		for job := range jobs {
			aiWG.Add(1)
			aiSem <- struct{}{}
			go func(job aiJob) {
				defer aiWG.Done()
				defer func() { <-aiSem }()
				if ctx.Err() != nil {
					return
				}
				promptTmpl := titlePrompt
				if job.mode != "title" {
					promptTmpl = articlePrompt
				}
				prompt := strings.ReplaceAll(promptTmpl, "{keyword}", job.row.Keyword)
				content, err := aiClient.Generate(ctx, prompt)
				if err != nil {
					log.Printf("  AI FAIL %s <%s>: %v", job.row.Keyword, job.mode, err)
					return
				}
				if job.mode == "title" {
					content = ai.CleanTitle(content)
					db.UpdateAITitle(job.writer, job.row.ID, content)
				} else {
					db.UpdateAIContent(job.writer, job.row.ID, content)
				}
				cur := atomic.AddInt64(&aiCounter, 1)
				if cur%100 == 0 {
					log.Printf("  AI progress: %d items done", cur)
				}
			}(job)
		}
	}()

	sem := make(chan struct{}, env.Cfg.Concurrency)
	var scrapeWG sync.WaitGroup
	var done atomic.Int64

	for _, item := range rows {
		if ctx.Err() != nil {
			break
		}
		scrapeWG.Add(1)
		sem <- struct{}{}
		go func(item db.WorkRow) {
			defer scrapeWG.Done()
			defer func() { <-sem }()
			if ctx.Err() != nil {
				return
			}
			results, err := env.Pin.Search(ctx, item.Keyword, env.Cfg.ImageResult, env.Cfg.MinSnippet)
			if err != nil {
				log.Printf("  FAIL %s: %v", item.Keyword, err)
				return
			}
			if len(results.Images) == 0 {
				log.Printf("  [done] %s -> 0 images (status stays 0)", item.Keyword)
				return
			}
			trimmed := results.Images
			if len(trimmed) > env.Cfg.ImageResult {
				trimmed = trimmed[:env.Cfg.ImageResult]
			}
			imagesJSON := marshalImages(trimmed)
			cleaned := make([]string, 0, len(results.Descriptions))
			for _, d := range results.Descriptions {
				if fd := util.FormatDescription(d); fd != "" {
					cleaned = append(cleaned, fd)
				}
			}
			snippet := mergeSnippet(item.Existing, nil, cleaned)
			db.UpdateImagesAndSnippet(ctx, writer, item.ID, imagesJSON, snippet)

			jobs <- aiJob{row: item, mode: "title", writer: writer}
			jobs <- aiJob{row: item, mode: "article", writer: writer}

			total := done.Add(1)
			if total%50 == 0 {
				log.Printf("  [%d/%d] scrape progress | %s", total, len(rows), env.Metrics.Snapshot())
			}
		}(item)
	}

	scrapeWG.Wait()
	close(jobs)
	consumers.Wait()
	aiWG.Wait()
	writer.Close()
	db.Close(conn)

	log.Printf("Completed %s: %d scraped, %d AI items", filepath.Base(dbPath), done.Load(), aiCounter)
	return aiCounter, nil
}
