package mongo

import (
	"context"
	"crypto/sha256"
	"encoding/hex"
	"errors"
	"fmt"
	"strings"
	"time"

	"go.mongodb.org/mongo-driver/bson"
	"go.mongodb.org/mongo-driver/mongo"
	"go.mongodb.org/mongo-driver/mongo/options"
)

const collectionIndexingRuns = "indexing_runs"
const collectionIndexingDiagnostics = "indexing_diagnostics"
const commitIntelFailureThreshold = 0.05

func resolveExpectedCommitDeltas(expectedCommits, pinnedDeltas int, completedCount, deltaCount int64) int64 {
	if pinnedDeltas > 0 {
		return int64(pinnedDeltas)
	}
	if expectedCommits <= 0 {
		return 0
	}
	if completedCount >= int64(expectedCommits) {
		return deltaCount
	}
	if deltaCount > 0 && completedCount >= deltaCount {
		skipped := completedCount - deltaCount
		remaining := int64(expectedCommits) - completedCount
		if remaining <= skipped {
			return deltaCount
		}
	}
	return int64(expectedCommits)
}

func commitIntelStageStatus(
	analysisCount, deltaCount, completedCount int64,
	expectedCommits, pinnedDeltas int,
) string {
	target := resolveExpectedCommitDeltas(expectedCommits, pinnedDeltas, completedCount, deltaCount)
	if expectedCommits <= 0 {
		return "completed"
	}
	if target > 0 && analysisCount >= target {
		return "completed"
	}
	if analysisCount == 1 {
		return "running"
	}
	return ""
}

// UpdateCommitIntelStage marks commit intelligence stage progress for a snapshot.
func (s *Store) UpdateCommitIntelStage(ctx context.Context, repoID, snapshotID string) error {
	if !s.indexingRunsEnabled {
		return nil
	}
	if err := ctx.Err(); err != nil {
		return err
	}

	col := s.indexingRunsCollection(collectionIndexingRuns)
	var run struct {
		ExpectedCommits      int `bson:"expected_commits"`
		ExpectedCommitDeltas int `bson:"expected_commit_deltas"`
	}
	err := col.FindOne(ctx, bson.M{
		"repo_id":     repoID,
		"snapshot_id": snapshotID,
	}).Decode(&run)
	if err != nil {
		if errors.Is(err, mongo.ErrNoDocuments) {
			return nil
		}
		return fmt.Errorf("load indexing run: %w", err)
	}

	analysisCount, err := s.CountAnalysesBySnapshot(ctx, repoID, snapshotID)
	if err != nil {
		return err
	}
	deltaCount, err := s.CountGraphDeltasBySnapshot(ctx, repoID, snapshotID)
	if err != nil {
		return err
	}
	completedCount, err := s.CountCompletedCommitsBySnapshot(ctx, repoID, snapshotID)
	if err != nil {
		return err
	}

	status := commitIntelStageStatus(
		analysisCount,
		deltaCount,
		completedCount,
		run.ExpectedCommits,
		run.ExpectedCommitDeltas,
	)
	if status == "" {
		return nil
	}

	_, err = col.UpdateOne(ctx, bson.M{
		"repo_id":     repoID,
		"snapshot_id": snapshotID,
	}, bson.M{"$set": bson.M{"stages.commit_intel": status}})
	if err != nil {
		return fmt.Errorf("update commit_intel stage: %w", err)
	}
	return nil
}

// IncrementFailedCommit records a permanent commit-intel failure exactly once per commit.
func (s *Store) IncrementFailedCommit(ctx context.Context, repoID, snapshotID, commitSHA, eventID, message string) error {
	if !s.indexingRunsEnabled {
		return nil
	}
	if err := ctx.Err(); err != nil {
		return err
	}

	col := s.indexingRunsCollection(collectionIndexingRuns)
	filter := bson.M{"repo_id": repoID, "snapshot_id": snapshotID}
	match := bson.M{
		"repo_id":            repoID,
		"snapshot_id":        snapshotID,
		"failed_commit_shas": bson.M{"$ne": commitSHA},
	}
	result, err := col.UpdateOne(ctx, match, bson.M{
		"$addToSet": bson.M{"failed_commit_shas": commitSHA},
		"$inc":      bson.M{"failed_commits": 1},
	})
	if err != nil {
		return fmt.Errorf("increment failed_commits: %w", err)
	}

	var run struct {
		ExpectedCommits      int `bson:"expected_commits"`
		ExpectedCommitDeltas int `bson:"expected_commit_deltas"`
		FailedCommits        int `bson:"failed_commits"`
	}
	if err := col.FindOne(ctx, filter).Decode(&run); err != nil {
		if errors.Is(err, mongo.ErrNoDocuments) {
			return nil
		}
		return fmt.Errorf("load indexing run after failed commit increment: %w", err)
	}

	target := run.ExpectedCommitDeltas
	if target <= 0 {
		target = run.ExpectedCommits
	}
	failureRate := failureRate(run.FailedCommits, target)
	set := bson.M{"commit_intel_failure_rate": failureRate}
	if result.MatchedCount > 0 && commitIntelFailureThreshold > 0 && failureRate >= commitIntelFailureThreshold {
		set["stages.commit_intel"] = "failed"
	}
	if _, err := col.UpdateOne(ctx, filter, bson.M{"$set": set}); err != nil {
		return fmt.Errorf("update commit_intel failure rate: %w", err)
	}

	return s.saveIndexingDiagnostic(ctx, repoID, snapshotID, "commit_intel", "commit", commitSHA, eventID, "commit_intel_failed", message)
}

func (s *Store) saveIndexingDiagnostic(ctx context.Context, repoID, snapshotID, stage, sourceType, sourceID, eventID, code, message string) error {
	id := diagnosticID(repoID, snapshotID, stage, sourceType, sourceID, eventID, code)
	doc := bson.M{
		"diagnostic_id": id,
		"repo_id":       repoID,
		"snapshot_id":   snapshotID,
		"stage":         stage,
		"source_type":   sourceType,
		"source_id":     sourceID,
		"event_id":      eventID,
		"code":          code,
		"severity":      "error",
		"message":       message,
		"created_at":    time.Now().UTC(),
	}
	_, err := s.indexingRunsCollection(collectionIndexingDiagnostics).UpdateOne(
		ctx,
		bson.M{"diagnostic_id": id},
		bson.M{"$set": doc},
		options.Update().SetUpsert(true),
	)
	if err != nil {
		return fmt.Errorf("save indexing diagnostic: %w", err)
	}
	return nil
}

func failureRate(failed, expected int) float64 {
	if failed <= 0 || expected <= 0 {
		return 0
	}
	return float64(failed) / float64(expected)
}

func diagnosticID(parts ...string) string {
	sum := sha256.Sum256([]byte(strings.Join(parts, "|")))
	return "diag_" + hex.EncodeToString(sum[:8])
}
