package mongo

import (
	"context"
	"errors"
	"fmt"
	"strings"
	"time"

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

	"bit.admedia.com/scm/ad/adpilot-indexing-code-parser.com/internal/contracts/graph"
	"bit.admedia.com/scm/ad/adpilot-indexing-code-parser.com/internal/store"
)

type snapshotGraphDoc struct {
	ArtifactID    string       `bson:"artifact_id"`
	RepoID        string       `bson:"repo_id"`
	SnapshotID    string       `bson:"snapshot_id"`
	CommitSHA     string       `bson:"commit_sha"`
	SchemaVersion string       `bson:"schema_version"`
	Chunked       bool         `bson:"chunked,omitempty"`
	PartCount     int          `bson:"part_count,omitempty"`
	NodeCount     int          `bson:"node_count"`
	EdgeCount     int          `bson:"edge_count"`
	Nodes         []graph.Node `bson:"nodes,omitempty"`
	Edges         []graph.Edge `bson:"edges,omitempty"`
	UpdatedAt     time.Time    `bson:"updated_at"`
}

// SaveSnapshotGraph upserts a merged snapshot graph artifact.
func (s *Store) SaveSnapshotGraph(ctx context.Context, artifact graph.Artifact, artifactID string) (string, string, error) {
	if err := ctx.Err(); err != nil {
		return "", "", err
	}
	if strings.TrimSpace(artifactID) == "" {
		artifactID = store.BuildMergedSnapshotArtifactID(artifact.SnapshotID)
	}
	if artifact.SchemaVersion == "" {
		artifact.SchemaVersion = graph.SchemaVersionV1
	}

	parts := splitGraphPayload(artifact.Nodes, artifact.Edges, s.maxPartBytes)
	if len(parts) == 1 {
		doc := snapshotGraphDoc{
			ArtifactID:    artifactID,
			RepoID:        artifact.RepoID,
			SnapshotID:    artifact.SnapshotID,
			CommitSHA:     artifact.CommitSHA,
			SchemaVersion: artifact.SchemaVersion,
			Chunked:       false,
			PartCount:     0,
			NodeCount:     len(artifact.Nodes),
			EdgeCount:     len(artifact.Edges),
			Nodes:         parts[0].Nodes,
			Edges:         parts[0].Edges,
			UpdatedAt:     time.Now().UTC(),
		}
		if err := s.upsertSnapshotGraphManifest(ctx, doc); err != nil {
			if mongoIsDocumentTooLarge(err) {
				return s.saveChunkedSnapshotGraph(ctx, artifactID, artifact)
			}
			return "", "", fmt.Errorf("upsert snapshot graph: %w", err)
		}
		if err := s.deleteGraphParts(ctx, collectionSnapshotGraphParts, artifactID); err != nil {
			return "", "", err
		}
		return artifactID, store.MongoURI(collectionSnapshotGraphs, artifactID), nil
	}

	return s.saveChunkedSnapshotGraph(ctx, artifactID, artifact)
}

func (s *Store) saveChunkedSnapshotGraph(ctx context.Context, artifactID string, artifact graph.Artifact) (string, string, error) {
	parts := splitGraphPayload(artifact.Nodes, artifact.Edges, s.maxPartBytes)
	if len(parts) == 1 {
		parts = splitGraphPayload(artifact.Nodes, artifact.Edges, s.maxPartBytes/2)
	}
	if len(parts) == 1 {
		return "", "", fmt.Errorf("upsert snapshot graph: graph payload exceeds mongodb document limit")
	}

	manifest := snapshotGraphDoc{
		ArtifactID:    artifactID,
		RepoID:        artifact.RepoID,
		SnapshotID:    artifact.SnapshotID,
		CommitSHA:     artifact.CommitSHA,
		SchemaVersion: artifact.SchemaVersion,
		Chunked:       true,
		PartCount:     len(parts),
		NodeCount:     len(artifact.Nodes),
		EdgeCount:     len(artifact.Edges),
		UpdatedAt:     time.Now().UTC(),
	}
	if err := s.upsertSnapshotGraphManifest(ctx, manifest); err != nil {
		return "", "", fmt.Errorf("upsert snapshot graph manifest: %w", err)
	}
	if err := s.saveGraphParts(ctx, collectionSnapshotGraphParts, artifactID, parts); err != nil {
		return "", "", fmt.Errorf("save snapshot graph parts: %w", err)
	}
	return artifactID, store.MongoURI(collectionSnapshotGraphs, artifactID), nil
}

func (s *Store) upsertSnapshotGraphManifest(ctx context.Context, doc snapshotGraphDoc) error {
	_, err := s.collection(collectionSnapshotGraphs).UpdateOne(
		ctx,
		bson.M{"artifact_id": doc.ArtifactID},
		bson.M{"$set": doc},
		options.Update().SetUpsert(true),
	)
	return err
}

// GetSnapshotGraphBySnapshot loads a merged snapshot graph when present.
func (s *Store) GetSnapshotGraphBySnapshot(ctx context.Context, repoID, snapshotID string) (graph.Artifact, string, bool, error) {
	var empty graph.Artifact
	if err := ctx.Err(); err != nil {
		return empty, "", false, err
	}

	var doc snapshotGraphDoc
	err := s.collection(collectionSnapshotGraphs).FindOne(ctx, bson.M{
		"repo_id":     repoID,
		"snapshot_id": snapshotID,
	}).Decode(&doc)
	if err != nil {
		if errors.Is(err, mongo.ErrNoDocuments) {
			return empty, "", false, nil
		}
		return empty, "", false, fmt.Errorf("find snapshot graph: %w", err)
	}

	artifact, err := s.artifactFromSnapshotDoc(ctx, doc)
	if err != nil {
		return empty, "", false, err
	}
	return artifact, doc.ArtifactID, true, nil
}

func (s *Store) artifactFromSnapshotDoc(ctx context.Context, doc snapshotGraphDoc) (graph.Artifact, error) {
	nodes, edges := doc.Nodes, doc.Edges
	if resolveChunked(doc.Chunked, doc.PartCount, len(doc.Nodes), len(doc.Edges)) {
		parts, err := s.loadGraphParts(ctx, collectionSnapshotGraphParts, doc.ArtifactID, doc.PartCount)
		if err != nil {
			return graph.Artifact{}, fmt.Errorf("load snapshot graph parts: %w", err)
		}
		nodes, edges = reassembleGraphPayload(parts)
	}

	return graph.Artifact{
		SchemaVersion: doc.SchemaVersion,
		RepoID:        doc.RepoID,
		CommitSHA:     doc.CommitSHA,
		SnapshotID:    doc.SnapshotID,
		Nodes:         nodes,
		Edges:         edges,
	}, nil
}
