package mongo

import (
	"fmt"

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

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

const (
	// mongoMaxDocumentBytes is MongoDB's per-document BSON size limit.
	mongoMaxDocumentBytes = 16 * 1024 * 1024
	// defaultMaxPartBytes leaves headroom below the MongoDB document limit.
	defaultMaxPartBytes = 10 * 1024 * 1024
)

// graphPayloadPart is one persisted chunk of nodes and/or edges.
type graphPayloadPart struct {
	Nodes []graph.Node
	Edges []graph.Edge
}

func effectiveMaxPartBytes(maxPartBytes int) int {
	if maxPartBytes <= 0 {
		return defaultMaxPartBytes
	}
	if maxPartBytes > mongoMaxDocumentBytes {
		return mongoMaxDocumentBytes - (256 * 1024)
	}
	return maxPartBytes
}

// splitGraphPayload splits nodes and edges into BSON-safe parts.
// Part 0 always contains all nodes plus the first edge batch; later parts hold edges only.
func splitGraphPayload(nodes []graph.Node, edges []graph.Edge, maxPartBytes int) []graphPayloadPart {
	maxPartBytes = effectiveMaxPartBytes(maxPartBytes)

	single := graphPayloadPart{Nodes: nodes, Edges: edges}
	if size, err := estimatePartBSONSize(single); err == nil && size <= maxPartBytes {
		return []graphPayloadPart{single}
	}

	parts := make([]graphPayloadPart, 0, 1+len(edges)/1000)
	current := graphPayloadPart{Nodes: nodes}
	includeNodes := true

	for _, edge := range edges {
		trial := graphPayloadPart{
			Nodes: current.Nodes,
			Edges: append(append([]graph.Edge(nil), current.Edges...), edge),
		}
		if !includeNodes {
			trial.Nodes = nil
		}

		size, err := estimatePartBSONSize(trial)
		if err != nil {
			size = maxPartBytes + 1
		}

		if size > maxPartBytes && len(current.Edges) > 0 {
			parts = append(parts, current)
			current = graphPayloadPart{Edges: []graph.Edge{edge}}
			includeNodes = false
			continue
		}

		current.Edges = append(current.Edges, edge)
		if !includeNodes {
			current.Nodes = nil
		}
	}

	if len(current.Nodes) > 0 || len(current.Edges) > 0 {
		parts = append(parts, current)
	}

	if len(parts) == 0 {
		return []graphPayloadPart{{Nodes: nodes}}
	}
	return parts
}

func reassembleGraphPayload(parts []graphPayloadPart) (nodes []graph.Node, edges []graph.Edge) {
	for _, part := range parts {
		if len(part.Nodes) > 0 {
			nodes = append(nodes, part.Nodes...)
		}
		if len(part.Edges) > 0 {
			edges = append(edges, part.Edges...)
		}
	}
	return nodes, edges
}

func estimatePartBSONSize(part graphPayloadPart) (int, error) {
	data, err := bson.Marshal(bson.M{
		"nodes": part.Nodes,
		"edges": part.Edges,
	})
	if err != nil {
		return 0, fmt.Errorf("marshal graph part: %w", err)
	}
	return len(data), nil
}
