package db

import (
	"context"
	"database/sql"
	"fmt"
	"log"
	"sync"
	"time"
)

type Update struct {
	SQL  string
	Args []interface{}
}

type Writer struct {
	conn       *sql.DB
	ch         chan Update
	done       chan struct{}
	wg         sync.WaitGroup
	batchSize  int
	flushEvery time.Duration

	mu      sync.Mutex
	total   int64
	batches int64
}

type WriterOptions struct {
	BatchSize  int
	FlushEvery time.Duration
}

func NewWriter(conn *sql.DB, opts WriterOptions) *Writer {
	if opts.BatchSize <= 0 {
		opts.BatchSize = 100
	}
	if opts.FlushEvery <= 0 {
		opts.FlushEvery = 2 * time.Second
	}
	w := &Writer{
		conn:       conn,
		ch:         make(chan Update, opts.BatchSize*4),
		done:       make(chan struct{}),
		batchSize:  opts.BatchSize,
		flushEvery: opts.FlushEvery,
	}
	w.wg.Add(1)
	go w.loop()
	return w
}

func (w *Writer) Submit(u Update) {
	w.ch <- u
}

func (w *Writer) loop() {
	defer w.wg.Done()
	batch := make([]Update, 0, w.batchSize)
	ticker := time.NewTicker(w.flushEvery)
	defer ticker.Stop()

	flush := func() {
		if len(batch) == 0 {
			return
		}
		if err := w.commit(batch); err != nil {
			log.Printf("db writer: batch commit failed (%d updates): %v", len(batch), err)
		}
		batch = batch[:0]
	}

	for {
		select {
		case u, ok := <-w.ch:
			if !ok {
				flush()
				close(w.done)
				return
			}
			batch = append(batch, u)
			if len(batch) >= w.batchSize {
				flush()
			}
		case <-ticker.C:
			flush()
		}
	}
}

// commit writes one batch in a single transaction. If the process dies before
// commit, the affected rows keep their previous (NULL) values and are picked up
// again on the next run: nothing is half-written.
func (w *Writer) commit(batch []Update) error {
	tx, err := w.conn.Begin()
	if err != nil {
		return fmt.Errorf("begin: %w", err)
	}
	for _, u := range batch {
		if _, err := tx.Exec(u.SQL, u.Args...); err != nil {
			tx.Rollback()
			return fmt.Errorf("exec %q: %w", u.SQL, err)
		}
	}
	if err := tx.Commit(); err != nil {
		tx.Rollback()
		return fmt.Errorf("commit: %w", err)
	}
	w.mu.Lock()
	w.total += int64(len(batch))
	w.batches++
	w.mu.Unlock()
	return nil
}

// Close flushes all pending updates and stops the writer goroutine.
func (w *Writer) Close() {
	close(w.ch)
	w.wg.Wait()
	<-w.done
}

func (w *Writer) Stats() (updates, batches int64) {
	w.mu.Lock()
	defer w.mu.Unlock()
	return w.total, w.batches
}

func UpdateImages(ctx context.Context, w *Writer, id int64, imagesJSON string) {
	w.Submit(Update{
		SQL:  `UPDATE posts SET images = ?, status = 1, updated_at = datetime('now') WHERE id = ?`,
		Args: []interface{}{imagesJSON, id},
	})
}

// UpdateImagesAndSnippet writes both scraped columns plus the status flip in
// one statement so the merged scrape costs a single UPDATE per row.
func UpdateImagesAndSnippet(ctx context.Context, w *Writer, id int64, imagesJSON, snippetJSON string) {
	w.Submit(Update{
		SQL:  `UPDATE posts SET images = ?, snippet = ?, status = 1, updated_at = datetime('now') WHERE id = ?`,
		Args: []interface{}{imagesJSON, snippetJSON, id},
	})
}

func UpdateSnippet(ctx context.Context, w *Writer, id int64, snippetJSON string) {
	w.Submit(Update{
		SQL:  `UPDATE posts SET snippet = ?, updated_at = datetime('now') WHERE id = ?`,
		Args: []interface{}{snippetJSON, id},
	})
}

func UpdateAITitle(w *Writer, id int64, title string) {
	w.Submit(Update{
		SQL:  `UPDATE posts SET ai_title = ?, updated_at = datetime('now') WHERE id = ?`,
		Args: []interface{}{title, id},
	})
}

func UpdateAIContent(w *Writer, id int64, content string) {
	w.Submit(Update{
		SQL:  `UPDATE posts SET ai_content = ?, updated_at = datetime('now') WHERE id = ?`,
		Args: []interface{}{content, id},
	})
}
