package mongo

import (
	"context"
	"fmt"
	"time"

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

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

const (
	collectionCodeGraphArtifactParts = "code_graph_artifact_parts"
	collectionSnapshotGraphParts     = "snapshot_graph_parts"
)

type graphPartDoc struct {
	ArtifactID string       `bson:"artifact_id"`
	PartIndex  int          `bson:"part_index"`
	Nodes      []graph.Node `bson:"nodes,omitempty"`
	Edges      []graph.Edge `bson:"edges,omitempty"`
	UpdatedAt  time.Time    `bson:"updated_at"`
}

func (s *Store) saveGraphParts(ctx context.Context, collection string, artifactID string, parts []graphPayloadPart) error {
	col := s.collection(collection)
	now := time.Now().UTC()

	for i, part := range parts {
		doc := graphPartDoc{
			ArtifactID: artifactID,
			PartIndex:  i,
			Nodes:      part.Nodes,
			Edges:      part.Edges,
			UpdatedAt:  now,
		}
		_, err := col.UpdateOne(
			ctx,
			bson.M{"artifact_id": artifactID, "part_index": i},
			bson.M{"$set": doc},
			options.Update().SetUpsert(true),
		)
		if err != nil {
			return fmt.Errorf("upsert graph part %d: %w", i, err)
		}
	}

	if err := s.deleteGraphPartsFromIndex(ctx, collection, artifactID, len(parts)); err != nil {
		return err
	}
	return nil
}

func (s *Store) loadGraphParts(ctx context.Context, collection, artifactID string, partCount int) ([]graphPayloadPart, error) {
	if partCount <= 0 {
		return nil, fmt.Errorf("part_count must be positive for chunked artifact %s", artifactID)
	}

	col := s.collection(collection)
	cursor, err := col.Find(ctx, bson.M{"artifact_id": artifactID}, options.Find().SetSort(bson.D{{Key: "part_index", Value: 1}}))
	if err != nil {
		return nil, fmt.Errorf("find graph parts: %w", err)
	}
	defer cursor.Close(ctx)

	parts := make([]graphPayloadPart, partCount)
	loaded := 0
	for cursor.Next(ctx) {
		var doc graphPartDoc
		if err := cursor.Decode(&doc); err != nil {
			return nil, fmt.Errorf("decode graph part: %w", err)
		}
		if doc.PartIndex < 0 || doc.PartIndex >= partCount {
			continue
		}
		parts[doc.PartIndex] = graphPayloadPart{Nodes: doc.Nodes, Edges: doc.Edges}
		loaded++
	}
	if err := cursor.Err(); err != nil {
		return nil, fmt.Errorf("iterate graph parts: %w", err)
	}
	if loaded != partCount {
		return nil, fmt.Errorf("expected %d graph parts for %s, found %d", partCount, artifactID, loaded)
	}
	return parts, nil
}

func (s *Store) deleteGraphParts(ctx context.Context, collection, artifactID string) error {
	_, err := s.collection(collection).DeleteMany(ctx, bson.M{"artifact_id": artifactID})
	if err != nil {
		return fmt.Errorf("delete graph parts for %s: %w", artifactID, err)
	}
	return nil
}

func (s *Store) deleteGraphPartsFromIndex(ctx context.Context, collection, artifactID string, fromIndex int) error {
	_, err := s.collection(collection).DeleteMany(ctx, bson.M{
		"artifact_id": artifactID,
		"part_index":  bson.M{"$gte": fromIndex},
	})
	if err != nil {
		return fmt.Errorf("delete trailing graph parts for %s: %w", artifactID, err)
	}
	return nil
}

func mongoIsDocumentTooLarge(err error) bool {
	if err == nil {
		return false
	}
	msg := err.Error()
	return containsIgnoreCase(msg, "document is too large") ||
		containsIgnoreCase(msg, "BSONObjectTooLarge") ||
		containsIgnoreCase(msg, "obj size")
}

// resolveChunked uses explicit chunked flag when set; otherwise infers from stored fields.
func resolveChunked(docChunked bool, partCount, nodeLen, edgeLen int) bool {
	if docChunked {
		return true
	}
	if partCount > 0 && nodeLen == 0 && edgeLen == 0 {
		return true
	}
	return false
}

func containsIgnoreCase(s, substr string) bool {
	return len(substr) > 0 && len(s) >= len(substr) &&
		(indexOfFold(s, substr) >= 0)
}

func indexOfFold(s, substr string) int {
	// simple case-insensitive contains
	lowerS := []byte(s)
	lowerSub := []byte(substr)
	for i := 0; i+len(lowerSub) <= len(lowerS); i++ {
		match := true
		for j := 0; j < len(lowerSub); j++ {
			a, b := lowerS[i+j], lowerSub[j]
			if a >= 'A' && a <= 'Z' {
				a += 'a' - 'A'
			}
			if b >= 'A' && b <= 'Z' {
				b += 'a' - 'A'
			}
			if a != b {
				match = false
				break
			}
		}
		if match {
			return i
		}
	}
	return -1
}
