package protocolgen

import (
	"encoding/json"
	"fmt"
	"go/format"
	"reflect"
	"strconv"
	"strings"

	"reasonix/internal/extension/protocol"
)

// The SDK types artifact (SDKTypesArtifactPath in generate.go) is produced
// from the same frozen registry and reflection walk as the JSON Schema:
// every wire DTO and enum reachable from the registry, plus the payload
// documents addressed indirectly through json.RawMessage fields and the
// structured error envelope.

// sdkEnumPrefixes maps every reachable string enum type to the constant-name
// prefix its values carry in the SDK's public API. ProviderErrorCode has
// none: its values already read as Go names ("provider_failed" →
// ProviderFailed).
var sdkEnumPrefixes = map[string]string{
	"InterceptEvent":    "Event",
	"InterceptDecision": "Decision",
	"UIHostKind":        "UIHost",
	"UISurfaceKind":     "UISurface",
	"UIRequestKind":     "UIRequest",
	"UIFieldKind":       "UIField",
	"UISeverity":        "UISeverity",
	"ProviderRole":      "ProviderRole",
	"ProviderChunkType": "Chunk",
	"ProviderErrorCode": "",
	"ContentEncoding":   "Content",
	"ErrorReason":       "Err",
}

// sdkEnumConstantExceptions pins constant names that mechanical mangling
// would render differently from the SDK's established public API.
var sdkEnumConstantExceptions = map[string]map[string]string{
	"ProviderChunkType": {"tool_call_args_delta": "ChunkToolCallDelta"},
}

// sdkInitialisms upper-cases identifier parts that read as acronyms.
var sdkInitialisms = map[string]string{
	"ui": "UI", "tui": "TUI", "acp": "ACP", "utf8": "UTF8",
}

var rawMessageType = reflect.TypeFor[json.RawMessage]()

// sdkTypeWalk is the deterministic first-visit record of every named wire
// type reachable from the frozen registry (plus the extra roots below).
type sdkTypeWalk struct {
	enums map[string][]string // frozen enum value sets, by type name
	order []reflect.Type      // named types in discovery order
	kinds map[reflect.Type]string
	seen  map[reflect.Type]bool
}

// walkSDKTypes reflection-walks the frozen registry exactly like the schema
// builder does — registry params and non-notification results — and adds the
// roots no registry DTO references by name: the host UI payload documents
// (carried inside json.RawMessage payload fields) and the structured error
// envelope ProtocolErrorData.
func walkSDKTypes() (*sdkTypeWalk, error) {
	w := &sdkTypeWalk{
		enums: protocol.EnumValues(),
		kinds: map[reflect.Type]string{},
		seen:  map[reflect.Type]bool{},
	}
	var roots []reflect.Type
	for _, spec := range protocol.Registry() {
		roots = append(roots, spec.ParamsType)
		if !spec.Notification() {
			roots = append(roots, spec.ResultType)
		}
	}
	roots = append(roots,
		reflect.TypeFor[protocol.UIStatusPayload](),
		reflect.TypeFor[protocol.UICardPayload](),
		reflect.TypeFor[protocol.UIFormPayload](),
		reflect.TypeFor[protocol.UINotificationPayload](),
		reflect.TypeFor[protocol.ProtocolErrorData](),
	)
	for _, root := range roots {
		if err := w.visit(root); err != nil {
			return nil, err
		}
	}
	return w, nil
}

func (w *sdkTypeWalk) visit(typ reflect.Type) error {
	for typ.Kind() == reflect.Pointer {
		typ = typ.Elem()
	}
	if typ == rawMessageType {
		return nil
	}
	switch typ.Kind() {
	case reflect.Struct:
		if typ.Name() == "" {
			return fmt.Errorf("anonymous struct %v is not a named wire DTO", typ)
		}
		if w.seen[typ] {
			return nil
		}
		w.seen[typ] = true
		w.order = append(w.order, typ)
		w.kinds[typ] = "struct"
		for i := range typ.NumField() {
			field := typ.Field(i)
			if field.PkgPath != "" {
				continue
			}
			if field.Anonymous {
				return fmt.Errorf("embedded field %v is not supported in wire DTOs", field.Type)
			}
			if err := w.visit(field.Type); err != nil {
				return err
			}
		}
		return nil
	case reflect.String:
		if typ.PkgPath() == "" {
			return nil // predeclared string
		}
		if _, ok := w.enums[typ.Name()]; !ok {
			return fmt.Errorf("named string type %s (%v) is not a frozen enum", typ.Name(), typ)
		}
		if w.seen[typ] {
			return nil
		}
		w.seen[typ] = true
		w.order = append(w.order, typ)
		w.kinds[typ] = "enum"
		return nil
	case reflect.Slice, reflect.Array:
		return w.visit(typ.Elem())
	case reflect.Map:
		if err := w.visit(typ.Key()); err != nil {
			return err
		}
		return w.visit(typ.Elem())
	}
	// Predeclared scalars and unconstrained interfaces carry no named types.
	return nil
}

// generateSDKTypesGo renders the SDK's DTO mirror from the frozen walk.
func generateSDKTypesGo() ([]byte, error) {
	walk, err := walkSDKTypes()
	if err != nil {
		return nil, err
	}
	var out strings.Builder
	out.WriteString("// Code generated by cmd/extension-protocol-gen; DO NOT EDIT.\n")
	out.WriteString("\n")
	out.WriteString("// Package extension: Extension Protocol v2 wire DTOs, enums, method\n")
	out.WriteString("// names, frozen limits, and the frozen error table, mirrored from the\n")
	out.WriteString("// host's internal/extension/protocol package. Behavior (validators,\n")
	out.WriteString("// error constructors, helpers) lives in the handwritten files.\n")
	out.WriteString("package extension\n")
	out.WriteString("\n")
	out.WriteString("import \"encoding/json\"\n\n")

	emitSDKIdentity(&out)
	emitSDKLimits(&out)
	if err := emitSDKMethods(&out); err != nil {
		return nil, err
	}
	emitSDKErrorTable(&out)
	for _, typ := range walk.order {
		var err error
		switch walk.kinds[typ] {
		case "enum":
			err = emitSDKEnum(&out, walk, typ)
		case "struct":
			err = emitSDKStruct(&out, typ)
		}
		if err != nil {
			return nil, err
		}
	}
	if err := guardSDKSurfacePayloads(walk); err != nil {
		return nil, err
	}
	formatted, err := format.Source([]byte(out.String()))
	if err != nil {
		return nil, fmt.Errorf("format sdk types source: %w", err)
	}
	return formatted, nil
}

func emitSDKIdentity(out *strings.Builder) {
	out.WriteString("// ProtocolID is the immutable identity string peers exchange during the\n")
	out.WriteString("// initialize handshake.\n")
	fmt.Fprintf(out, "const ProtocolID = %q\n\n", protocol.ProtocolID)
	out.WriteString("// ProtocolMajor is the frozen major version of this protocol build.\n")
	fmt.Fprintf(out, "const ProtocolMajor = %d\n\n", protocol.ProtocolMajor)
	out.WriteString("// ProtocolVersion is the wire string form of ProtocolMajor carried in the\n")
	out.WriteString("// initialize handshake.\n")
	fmt.Fprintf(out, "const ProtocolVersion = %q\n\n", protocol.ProtocolVersion)
}

func emitSDKLimits(out *strings.Builder) {
	limits := protocol.FrozenLimits()
	out.WriteString("// Frozen wire limits. These constants are part of the protocol contract.\n")
	out.WriteString("const (\n")
	out.WriteString("\t// FrameBytes caps one JSON-RPC frame on the extension transport.\n")
	fmt.Fprintf(out, "\tFrameBytes = %d\n", limits.FrameBytes)
	out.WriteString("\t// ExternalizeFieldBytes is the threshold above which an externalizable\n")
	out.WriteString("\t// payload must move into a content ref instead of traveling inline.\n")
	fmt.Fprintf(out, "\tExternalizeFieldBytes = %d\n", limits.ExternalizeFieldBytes)
	out.WriteString("\t// ContentRefChunkBytes caps one host/content/read chunk.\n")
	fmt.Fprintf(out, "\tContentRefChunkBytes = %d\n", limits.ContentRefChunkBytes)
	out.WriteString("\t// ContentRefObjectBytes caps one externalized object.\n")
	fmt.Fprintf(out, "\tContentRefObjectBytes = %d\n", limits.ContentRefObjectBytes)
	out.WriteString(")\n\n")
}

func emitSDKMethods(out *strings.Builder) error {
	out.WriteString("// Method names, frozen for Extension Protocol v2.\n")
	out.WriteString("const (\n")
	seen := map[string]bool{}
	for _, spec := range protocol.Registry() {
		name := sdkMethodConstantName(string(spec.Name))
		if seen[name] {
			return fmt.Errorf("method constant name collision: %s", name)
		}
		seen[name] = true
		fmt.Fprintf(out, "\t%s = %q\n", name, string(spec.Name))
	}
	out.WriteString(")\n\n")
	return nil
}

func emitSDKErrorTable(out *strings.Builder) {
	out.WriteString("// DomainErrorCode is the JSON-RPC code every extension domain error uses on\n")
	out.WriteString("// the wire. The structured ProtocolErrorData reason distinguishes them.\n")
	fmt.Fprintf(out, "const DomainErrorCode = %d\n\n", protocol.DomainErrorCode)
	out.WriteString("// errorSpec is one frozen error table entry: the JSON-RPC code, the\n")
	out.WriteString("// generic wire message, and whether the call may be retried.\n")
	out.WriteString("type errorSpec struct {\n\tCode int\n\tMessage string\n\tRetryable bool\n}\n\n")
	out.WriteString("// frozenErrorSpecs mirrors the host's frozen error table. Adding an entry\n")
	out.WriteString("// is a conscious protocol change.\n")
	out.WriteString("var frozenErrorSpecs = map[ErrorReason]errorSpec{\n")
	for _, contract := range protocol.ErrorContracts() {
		fmt.Fprintf(out, "\t%s: {%s, %q, %t},\n",
			sdkEnumConstantName("ErrorReason", string(contract.Reason)),
			sdkErrorCodeName(contract.JSONRPCCode), contract.Message, contract.Retryable)
	}
	out.WriteString("}\n\n")
}

// sdkErrorCodeName renders a frozen JSON-RPC code with the SDK's symbolic
// constant where one exists (the standard codes live in wire.go).
func sdkErrorCodeName(code int) string {
	switch code {
	case -32600:
		return "CodeInvalidRequest"
	case -32601:
		return "CodeMethodNotFound"
	case -32602:
		return "CodeInvalidParams"
	case -32603:
		return "CodeInternal"
	case protocol.DomainErrorCode:
		return "DomainErrorCode"
	default:
		return strconv.Itoa(code)
	}
}

func emitSDKEnum(out *strings.Builder, walk *sdkTypeWalk, typ reflect.Type) error {
	name := typ.Name()
	if _, ok := sdkEnumPrefixes[name]; !ok {
		return fmt.Errorf("enum %s has no constant prefix registered", name)
	}
	fmt.Fprintf(out, "// %s is a generated Extension Protocol v2 string enum.\n", name)
	fmt.Fprintf(out, "type %s string\n\n", name)
	out.WriteString("const (\n")
	seen := map[string]bool{}
	for _, value := range walk.enums[name] {
		constant := sdkEnumConstantName(name, value)
		if seen[constant] {
			return fmt.Errorf("enum constant name collision: %s", constant)
		}
		seen[constant] = true
		fmt.Fprintf(out, "\t%s %s = %q\n", constant, name, value)
	}
	out.WriteString(")\n\n")
	return nil
}

func emitSDKStruct(out *strings.Builder, typ reflect.Type) error {
	fmt.Fprintf(out, "// %s is a generated Extension Protocol v2 wire DTO.\n", typ.Name())
	fields := 0
	var body strings.Builder
	for i := range typ.NumField() {
		field := typ.Field(i)
		if field.PkgPath != "" {
			continue
		}
		rendered, err := renderSDKType(field.Type)
		if err != nil {
			return fmt.Errorf("%s field %s: %w", typ.Name(), field.Name, err)
		}
		fmt.Fprintf(&body, "\t%s %s `%s`\n", field.Name, rendered, string(field.Tag))
		fields++
	}
	if fields == 0 {
		fmt.Fprintf(out, "type %s struct{}\n\n", typ.Name())
		return nil
	}
	fmt.Fprintf(out, "type %s struct {\n%s}\n\n", typ.Name(), body.String())
	return nil
}

// renderSDKType renders a field type as Go source. Named types resolve to
// their bare name — every named type reachable from the walk is either one of
// the mirrored DTOs/enums or encoding/json's RawMessage.
func renderSDKType(typ reflect.Type) (string, error) {
	if typ == rawMessageType {
		return "json.RawMessage", nil
	}
	switch typ.Kind() {
	case reflect.Pointer:
		elem, err := renderSDKType(typ.Elem())
		return "*" + elem, err
	case reflect.Slice:
		elem, err := renderSDKType(typ.Elem())
		return "[]" + elem, err
	case reflect.Map:
		key, err := renderSDKType(typ.Key())
		if err != nil {
			return "", err
		}
		elem, err := renderSDKType(typ.Elem())
		return "map[" + key + "]" + elem, err
	case reflect.Interface:
		if typ.NumMethod() == 0 {
			return "any", nil
		}
	case reflect.Struct:
		if typ.Name() != "" {
			return typ.Name(), nil
		}
	case reflect.String:
		if typ.Name() != "" {
			return typ.Name(), nil
		}
		return "string", nil
	case reflect.Bool:
		return "bool", nil
	case reflect.Int:
		return "int", nil
	case reflect.Int64:
		return "int64", nil
	case reflect.Uint64:
		return "uint64", nil
	case reflect.Float64:
		return "float64", nil
	}
	return "", fmt.Errorf("unsupported wire type %v", typ)
}

// guardSDKSurfacePayloads pins the surface-kind → payload-DTO convention: a
// new UISurfaceKind without a matching UI<Kind>Payload type in the generated
// set fails generation loudly instead of silently shipping an SDK that
// cannot build the new surface.
func guardSDKSurfacePayloads(walk *sdkTypeWalk) error {
	emitted := map[string]bool{}
	for _, typ := range walk.order {
		emitted[typ.Name()] = true
	}
	for _, kind := range walk.enums["UISurfaceKind"] {
		want := "UI" + sdkIdentifierFromValue(kind) + "Payload"
		if !emitted[want] {
			return fmt.Errorf("surface kind %q has no payload DTO %s in the generated set", kind, want)
		}
	}
	return nil
}

// sdkMethodConstantName mangles a wire method name into its Go constant
// name: "extension/ui/action" → MethodExtensionUIAction.
func sdkMethodConstantName(method string) string {
	return "Method" + sdkIdentifierFromValue(method)
}

// sdkEnumConstantName mangles an enum value into its Go constant name with
// the type's registered prefix; exceptions pin the established public API.
func sdkEnumConstantName(typeName, value string) string {
	if exceptions, ok := sdkEnumConstantExceptions[typeName]; ok {
		if name, ok := exceptions[value]; ok {
			return name
		}
	}
	return sdkEnumPrefixes[typeName] + sdkIdentifierFromValue(value)
}

// sdkIdentifierFromValue turns a wire value into Go identifier parts:
// "agent.before_start" → "AgentBeforeStart".
func sdkIdentifierFromValue(value string) string {
	parts := strings.FieldsFunc(value, func(r rune) bool {
		return r == '_' || r == '.' || r == '-' || r == '/'
	})
	var out strings.Builder
	for _, part := range parts {
		if initialism, ok := sdkInitialisms[part]; ok {
			out.WriteString(initialism)
			continue
		}
		out.WriteString(strings.ToUpper(part[:1]) + part[1:])
	}
	return out.String()
}
