package mongo

import (
	"context"
	"errors"
	"fmt"

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

// CountCompletedCommitsBySnapshot counts commits marked completed for ordered fan-in.
func (s *Store) CountCompletedCommitsBySnapshot(ctx context.Context, repoID, snapshotID string) (int64, error) {
	if err := ctx.Err(); err != nil {
		return 0, err
	}
	return s.collection(collectionCommitProcessingState).CountDocuments(ctx, bson.M{
		"repo_id":     repoID,
		"snapshot_id": snapshotID,
		"status":      commitStatusCompleted,
	})
}

// CountGraphDeltasBySnapshot counts persisted graph deltas for a snapshot.
func (s *Store) CountGraphDeltasBySnapshot(ctx context.Context, repoID, snapshotID string) (int64, error) {
	if err := ctx.Err(); err != nil {
		return 0, err
	}
	return s.collection(collectionGraphDeltas).CountDocuments(ctx, bson.M{
		"repo_id":     repoID,
		"snapshot_id": snapshotID,
	})
}

// ResolveExpectedCommitDeltas returns the number of delta-backed commits downstream
// stages should wait for. Skipped commits (no indexable code changes) are excluded.
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)
}

// CommitFanInComplete reports whether code-parser has finished commits.changed fan-in.
func CommitFanInComplete(expectedCommits, pinnedDeltas int, completedCount, deltaCount int64) bool {
	if expectedCommits <= 0 {
		return true
	}
	if pinnedDeltas > 0 {
		return true
	}
	if completedCount >= int64(expectedCommits) {
		return true
	}
	if deltaCount > 0 && completedCount >= deltaCount {
		skipped := completedCount - deltaCount
		remaining := int64(expectedCommits) - completedCount
		return remaining <= skipped
	}
	return false
}

// SyncCommitFanIn pins expected_commit_deltas once commit fan-in is complete.
func (s *Store) SyncCommitFanIn(ctx context.Context, repoID, snapshotID string) error {
	if !s.indexingRunsEnabled {
		return nil
	}
	if err := ctx.Err(); err != nil {
		return err
	}

	filter := bson.M{
		"repo_id":     repoID,
		"snapshot_id": snapshotID,
	}

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

	completedCount, err := s.CountCompletedCommitsBySnapshot(ctx, repoID, snapshotID)
	if err != nil {
		return err
	}
	deltaCount, err := s.CountGraphDeltasBySnapshot(ctx, repoID, snapshotID)
	if err != nil {
		return err
	}
	if !CommitFanInComplete(run.ExpectedCommits, run.ExpectedCommitDeltas, completedCount, deltaCount) {
		return nil
	}

	target := ResolveExpectedCommitDeltas(run.ExpectedCommits, run.ExpectedCommitDeltas, completedCount, deltaCount)
	_, err = col.UpdateOne(ctx, filter, bson.M{
		"$set": bson.M{"expected_commit_deltas": int(target)},
	})
	if err != nil {
		return fmt.Errorf("pin expected_commit_deltas: %w", err)
	}
	return nil
}
