package cache

import (
	"context"
	"crypto/sha256"
	"database/sql"
	"encoding/hex"
	"sort"
	"strconv"
	"strings"
	"time"
)

func (s *SQLiteStore) UpsertSourceGraph(ctx context.Context, graph SourceGraph) (err error) {
	tx, err := s.db.BeginTx(ctx, nil)
	if err != nil {
		return err
	}
	defer txRollbackOnError(tx, &err)
	if err = upsertSourceTx(ctx, tx, graph.Source); err != nil {
		return err
	}
	if err = upsertSearchProjectionTx(ctx, tx, graph.Source, s.useFTS); err != nil {
		return err
	}
	for _, identity := range graph.Identities {
		if identity.RepoID == "" {
			identity.RepoID = graph.Source.RepoID
		}
		if identity.SourceID == "" {
			identity.SourceID = graph.Source.ID
		}
		if err = upsertIdentityTx(ctx, tx, identity); err != nil {
			return err
		}
	}
	for _, link := range graph.Links {
		if link.RepoID == "" {
			link.RepoID = graph.Source.RepoID
		}
		if link.SourceID == "" {
			link.SourceID = graph.Source.ID
		}
		if err = upsertLinkTx(ctx, tx, link); err != nil {
			return err
		}
	}
	for _, review := range graph.PRReviewComments {
		if review.RepoID == "" {
			review.RepoID = graph.Source.RepoID
		}
		if review.SourceID == "" {
			review.SourceID = graph.Source.ID
		}
		if err = upsertPRReviewCommentTx(ctx, tx, review); err != nil {
			return err
		}
	}
	for _, discussion := range graph.PRReviewDiscussions {
		if discussion.RepoID == "" {
			discussion.RepoID = graph.Source.RepoID
		}
		if err = upsertPRReviewDiscussionTx(ctx, tx, discussion); err != nil {
			return err
		}
	}
	for _, position := range graph.PRReviewPositions {
		if position.RepoID == "" {
			position.RepoID = graph.Source.RepoID
		}
		if err = upsertPRReviewPositionTx(ctx, tx, position); err != nil {
			return err
		}
	}
	if graph.ReplaceChunks {
		if err = reconcileSourceChunksTx(ctx, tx, graph.Source.RepoID, graph.Source.ID, graph.Chunks); err != nil {
			return err
		}
	}
	for _, chunk := range graph.Chunks {
		if chunk.RepoID == "" {
			chunk.RepoID = graph.Source.RepoID
		}
		if chunk.SourceID == "" {
			chunk.SourceID = graph.Source.ID
		}
		if _, err = upsertChunkTx(ctx, tx, chunk); err != nil {
			return err
		}
	}
	if graph.SyncStatus != nil {
		status := *graph.SyncStatus
		if status.RepoID == "" {
			status.RepoID = graph.Source.RepoID
		}
		if status.SourceID == "" {
			status.SourceID = graph.Source.ID
		}
		if err = upsertSyncStatusTx(ctx, tx, status); err != nil {
			return err
		}
	}
	for _, event := range graph.SyncEvents {
		if event.RepoID == "" {
			event.RepoID = graph.Source.RepoID
		}
		if event.SourceID == "" {
			event.SourceID = graph.Source.ID
		}
		if err = recordSyncEventTx(ctx, tx, event); err != nil {
			return err
		}
	}
	for _, conflict := range graph.Conflicts {
		if conflict.RepoID == "" {
			conflict.RepoID = graph.Source.RepoID
		}
		if conflict.SourceID == "" {
			conflict.SourceID = graph.Source.ID
		}
		if err = upsertConflictTx(ctx, tx, conflict); err != nil {
			return err
		}
	}
	return tx.Commit()
}

func reconcileSourceChunksTx(ctx context.Context, tx *sql.Tx, repoID, sourceID string, replacements []Chunk) error {
	wanted := make(map[string]string, len(replacements))
	for _, chunk := range replacements {
		wanted[chunk.ID] = chunk.ContentHash
	}
	rows, err := tx.QueryContext(ctx, `SELECT id, content_hash FROM chunks WHERE repo_id = ? AND source_id = ?`, repoID, sourceID)
	if err != nil {
		return err
	}
	stale := []string{}
	for rows.Next() {
		var id, contentHash string
		if err := rows.Scan(&id, &contentHash); err != nil {
			rows.Close()
			return err
		}
		if replacementHash, ok := wanted[id]; !ok || replacementHash != contentHash {
			stale = append(stale, id)
		}
	}
	if err := rows.Err(); err != nil {
		rows.Close()
		return err
	}
	if err := rows.Close(); err != nil {
		return err
	}
	for _, id := range stale {
		if err := execTx(ctx, tx, `DELETE FROM chunks WHERE repo_id = ? AND id = ?`, repoID, id); err != nil {
			return err
		}
	}
	return nil
}

func (s *SQLiteStore) UpsertSource(ctx context.Context, source Source) (err error) {
	tx, err := s.db.BeginTx(ctx, nil)
	if err != nil {
		return err
	}
	defer txRollbackOnError(tx, &err)
	if err = upsertSourceTx(ctx, tx, source); err != nil {
		return err
	}
	if err = upsertSearchProjectionTx(ctx, tx, source, s.useFTS); err != nil {
		return err
	}
	return tx.Commit()
}

func upsertSourceTx(ctx context.Context, tx *sql.Tx, source Source) error {
	labels, err := marshalJSON(source.Labels)
	if err != nil {
		return err
	}
	if source.Provenance == "" {
		source.Provenance = ProvenanceFixture
	}
	createdAt := source.CreatedAt
	updatedAt := source.UpdatedAt
	if createdAt.IsZero() {
		createdAt = time.Unix(0, 0).UTC()
	}
	if updatedAt.IsZero() {
		updatedAt = createdAt
	}
	return execTx(ctx, tx, `INSERT INTO sources (repo_id, id, kind, path, title, body, status, labels, content_hash, provenance, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(repo_id, id) DO UPDATE SET kind = excluded.kind, path = excluded.path, title = excluded.title, body = excluded.body, status = excluded.status, labels = excluded.labels, content_hash = excluded.content_hash, provenance = excluded.provenance, updated_at = excluded.updated_at`,
		source.RepoID, source.ID, source.Kind, source.Path, source.Title, source.Body, source.Status, labels, source.ContentHash, string(source.Provenance), createdAt.Format(time.RFC3339Nano), updatedAt.Format(time.RFC3339Nano))
}

func upsertSearchProjectionTx(ctx context.Context, tx *sql.Tx, source Source, useFTS bool) error {
	if !useFTS {
		return nil
	}
	if err := execTx(ctx, tx, `DELETE FROM fts_index WHERE repo_id = ? AND source_id = ?`, source.RepoID, source.ID); err != nil {
		return err
	}
	return execTx(ctx, tx, `INSERT INTO fts_index (repo_id, source_id, path, title, body) VALUES (?, ?, ?, ?, ?)`, source.RepoID, source.ID, source.Path, source.Title, source.Body)
}

func (s *SQLiteStore) GetSource(ctx context.Context, id string) (Source, error) {
	source, err := s.scanSource(ctx, `SELECT repo_id, id, kind, path, title, body, status, labels, content_hash, provenance, created_at, updated_at FROM sources WHERE id = ? ORDER BY repo_id LIMIT 1`, id)
	if err != nil {
		return Source{}, err
	}
	aliases, err := s.GetIdentityMapScoped(ctx, source.RepoID, id)
	if err != nil {
		return Source{}, err
	}
	source.Aliases = aliases
	return source, nil
}

func (s *SQLiteStore) GetSourceScoped(ctx context.Context, repoID, id string) (Source, error) {
	source, err := s.scanSource(ctx, `SELECT repo_id, id, kind, path, title, body, status, labels, content_hash, provenance, created_at, updated_at FROM sources WHERE repo_id = ? AND id = ?`, repoID, id)
	if err != nil {
		return Source{}, err
	}
	aliases, err := s.GetIdentityMapScoped(ctx, repoID, id)
	if err != nil {
		return Source{}, err
	}
	source.Aliases = aliases
	return source, nil
}

func (s *SQLiteStore) ListSources(ctx context.Context, filter SourceFilter) ([]Source, error) {
	query := `SELECT repo_id, id, kind, path, title, body, status, labels, content_hash, provenance, created_at, updated_at FROM sources WHERE (? = '' OR repo_id = ?) AND (? = '' OR kind = ?) AND (? = '' OR status = ?) AND (? = '' OR provenance = ?) ORDER BY repo_id, id`
	args := []any{filter.RepoID, filter.RepoID, filter.Kind, filter.Kind, filter.Status, filter.Status, filter.Provenance, filter.Provenance}
	if filter.Limit > 0 {
		query += ` LIMIT ?`
		args = append(args, filter.Limit)
	}
	rows, err := s.db.QueryContext(ctx, query, args...)
	if err != nil {
		return nil, err
	}
	sources, err := scanSources(rows)
	closeErr := rows.Close()
	if err != nil {
		return nil, err
	}
	if closeErr != nil {
		return nil, closeErr
	}
	if err := s.attachSourceAliases(ctx, sources); err != nil {
		return nil, err
	}
	return sources, nil
}

// attachSourceAliases hydrates identities for a source list in bounded batches.
// Keeping this at the cache boundary prevents callers from falling into an N+1
// query pattern when they need stable and provider identities for list results.
func (s *SQLiteStore) attachSourceAliases(ctx context.Context, sources []Source) error {
	const batchSize = 500
	idsByRepo := make(map[string][]string)
	seen := make(map[string]struct{})
	for _, source := range sources {
		key := source.RepoID + "\x00" + source.ID
		if _, ok := seen[key]; ok {
			continue
		}
		seen[key] = struct{}{}
		idsByRepo[source.RepoID] = append(idsByRepo[source.RepoID], source.ID)
	}

	aliasesBySource := make(map[string][]Identity)
	for repoID, ids := range idsByRepo {
		for start := 0; start < len(ids); start += batchSize {
			end := start + batchSize
			if end > len(ids) {
				end = len(ids)
			}
			batch := ids[start:end]
			placeholders := strings.TrimSuffix(strings.Repeat("?,", len(batch)), ",")
			args := make([]any, 0, len(batch)+1)
			args = append(args, repoID)
			for _, id := range batch {
				args = append(args, id)
			}
			rows, err := s.db.QueryContext(ctx, `SELECT repo_id, source_id, alias_type, alias, remote_type, remote_id FROM identity_map WHERE repo_id = ? AND source_id IN (`+placeholders+`) ORDER BY source_id, alias_type, alias`, args...)
			if err != nil {
				return err
			}
			identities, scanErr := scanIdentities(rows)
			closeErr := rows.Close()
			if scanErr != nil {
				return scanErr
			}
			if closeErr != nil {
				return closeErr
			}
			for _, identity := range identities {
				key := identity.RepoID + "\x00" + identity.SourceID
				aliasesBySource[key] = append(aliasesBySource[key], identity)
			}
		}
	}
	for i := range sources {
		sources[i].Aliases = aliasesBySource[sources[i].RepoID+"\x00"+sources[i].ID]
	}
	return nil
}

func (s *SQLiteStore) scanSource(ctx context.Context, query string, args ...any) (Source, error) {
	rows, err := s.db.QueryContext(ctx, query, args...)
	if err != nil {
		return Source{}, err
	}
	defer rows.Close()
	sources, err := scanSources(rows)
	if err != nil {
		return Source{}, err
	}
	if len(sources) == 0 {
		return Source{}, notFoundErr("source", "")
	}
	return sources[0], nil
}

func scanSources(rows *sql.Rows) ([]Source, error) {
	var sources []Source
	for rows.Next() {
		var source Source
		var labelsRaw, provenance, createdRaw, updatedRaw string
		if err := rows.Scan(&source.RepoID, &source.ID, &source.Kind, &source.Path, &source.Title, &source.Body, &source.Status, &labelsRaw, &source.ContentHash, &provenance, &createdRaw, &updatedRaw); err != nil {
			return nil, err
		}
		labels, err := unmarshalJSON[[]string](labelsRaw)
		if err != nil {
			return nil, err
		}
		source.Labels = labels
		source.Provenance = Provenance(provenance)
		source.CreatedAt, _ = time.Parse(time.RFC3339Nano, createdRaw)
		source.UpdatedAt, _ = time.Parse(time.RFC3339Nano, updatedRaw)
		sources = append(sources, source)
	}
	return sources, rows.Err()
}

func (s *SQLiteStore) SearchSources(ctx context.Context, query SearchQuery) ([]SearchResult, error) {
	if s.useFTS {
		return s.searchSourcesFTS(ctx, query)
	}
	return s.searchSourcesFallback(ctx, query)
}

func (s *SQLiteStore) searchSourcesFallback(ctx context.Context, query SearchQuery) ([]SearchResult, error) {
	needle := normalizeSearchQuery(query.Query)
	rows, err := s.db.QueryContext(ctx, `SELECT repo_id, id, path, title, body, provenance FROM sources WHERE (? = '' OR repo_id = ?) AND (? = '' OR kind = ?) AND (? = '' OR provenance = ?) AND (lower(title) LIKE ? OR lower(body) LIKE ?) ORDER BY repo_id, id, path`, query.RepoID, query.RepoID, query.Kind, query.Kind, query.Provenance, query.Provenance, "%"+needle+"%", "%"+needle+"%")
	if err != nil {
		return nil, err
	}
	defer rows.Close()
	return scanSearchResults(rows, needle, query.Limit)
}

func (s *SQLiteStore) searchSourcesFTS(ctx context.Context, query SearchQuery) ([]SearchResult, error) {
	if err := s.repairSearchProjection(ctx, query.RepoID); err != nil {
		return nil, err
	}
	needle := normalizeSearchQuery(query.Query)
	match := ftsMatchQuery(needle)
	rows, err := s.db.QueryContext(ctx, `SELECT s.repo_id, s.id, s.path, s.title, s.body, s.provenance
FROM fts_index f
JOIN sources s ON s.repo_id = f.repo_id AND s.id = f.source_id
WHERE (? = '' OR s.repo_id = ?) AND (? = '' OR s.kind = ?) AND (? = '' OR s.provenance = ?) AND fts_index MATCH ?
ORDER BY s.repo_id, s.id, s.path`, query.RepoID, query.RepoID, query.Kind, query.Kind, query.Provenance, query.Provenance, match)
	if err != nil {
		return nil, err
	}
	defer rows.Close()
	return scanSearchResults(rows, needle, query.Limit)
}

func (s *SQLiteStore) repairSearchProjection(ctx context.Context, repoID string) (err error) {
	var missing int
	if err := s.db.QueryRowContext(ctx, `SELECT count(*)
FROM sources s
WHERE (? = '' OR s.repo_id = ?)
  AND NOT EXISTS (SELECT 1 FROM fts_index f WHERE f.repo_id = s.repo_id AND f.source_id = s.id)`, repoID, repoID).Scan(&missing); err != nil {
		return err
	}
	if missing == 0 {
		return nil
	}
	tx, err := s.db.BeginTx(ctx, nil)
	if err != nil {
		return err
	}
	defer txRollbackOnError(tx, &err)
	rows, err := tx.QueryContext(ctx, `SELECT repo_id, id, path, title, body FROM sources s
WHERE (? = '' OR s.repo_id = ?)
  AND NOT EXISTS (SELECT 1 FROM fts_index f WHERE f.repo_id = s.repo_id AND f.source_id = s.id)`, repoID, repoID)
	if err != nil {
		return err
	}
	defer rows.Close()
	for rows.Next() {
		var source Source
		if err = rows.Scan(&source.RepoID, &source.ID, &source.Path, &source.Title, &source.Body); err != nil {
			return err
		}
		if err = upsertSearchProjectionTx(ctx, tx, source, true); err != nil {
			return err
		}
	}
	if err = rows.Err(); err != nil {
		return err
	}
	return tx.Commit()
}

func scanSearchResults(rows *sql.Rows, needle string, limit int) ([]SearchResult, error) {
	var results []SearchResult
	for rows.Next() {
		var repoID, id, path, title, body, provenance string
		if err := rows.Scan(&repoID, &id, &path, &title, &body, &provenance); err != nil {
			return nil, err
		}
		results = append(results, SearchResult{RepoID: repoID, ID: id, Path: path, Title: title, Snippet: snippet(title+"\n"+body, needle), Score: searchScore(title, body, needle), Line: lineFor(body, needle), Provenance: Provenance(provenance)})
	}
	if err := rows.Err(); err != nil {
		return nil, err
	}
	sort.SliceStable(results, func(i, j int) bool {
		if results[i].Score != results[j].Score {
			return results[i].Score > results[j].Score
		}
		if results[i].ID != results[j].ID {
			return results[i].ID < results[j].ID
		}
		return results[i].Path < results[j].Path
	})
	if limit > 0 && len(results) > limit {
		results = results[:limit]
	}
	return results, nil
}

func normalizeSearchQuery(query string) string {
	return strings.ToLower(strings.TrimSpace(query))
}

func ftsMatchQuery(query string) string {
	parts := strings.Fields(query)
	if len(parts) == 0 {
		return `""`
	}
	for i, part := range parts {
		parts[i] = `"` + strings.ReplaceAll(part, `"`, `""`) + `"`
	}
	return strings.Join(parts, " AND ")
}

func searchScore(title, body, needle string) float64 {
	titleLower := strings.ToLower(title)
	bodyLower := strings.ToLower(body)
	score := 0.0
	if strings.Contains(titleLower, needle) {
		score += 10
	}
	score += float64(strings.Count(titleLower, needle) + strings.Count(bodyLower, needle))
	return score
}

func snippet(text, needle string) string {
	lower := strings.ToLower(text)
	idx := strings.Index(lower, needle)
	if idx < 0 {
		if len(text) > 80 {
			return text[:80]
		}
		return text
	}
	start := idx - 30
	if start < 0 {
		start = 0
	}
	end := idx + len(needle) + 30
	if end > len(text) {
		end = len(text)
	}
	return text[start:end]
}

func lineFor(body, needle string) int {
	idx := strings.Index(strings.ToLower(body), needle)
	if idx < 0 {
		return 1
	}
	return strings.Count(body[:idx], "\n") + 1
}

func deterministicChunkID(chunk Chunk) string {
	sum := sha256.Sum256([]byte(chunk.SourceID + "\x00" + chunk.ContentHash + "\x00" + strconv.Itoa(chunk.ByteStart)))
	return hex.EncodeToString(sum[:])
}