package main

import (
	"bytes"
	"crypto/rand"
	"crypto/sha256"
	"encoding/base64"
	"errors"
	"net/url"
	"strings"
	"testing"
	"time"

	"github.com/gorilla/websocket"
	"golang.org/x/crypto/nacl/box"
	"golang.org/x/crypto/nacl/sign"
)

func TestParseAuthResultRequiresExactCapabilities(t *testing.T) {
	valid := map[string]any{
		"status":       "done",
		"tmp_id":       "tmp-1",
		"app_id":       "app-1",
		"safe_id":      "safe-1",
		"capabilities": []any{"sso", "app_data"},
		"ver":          "1",
	}
	options := HandshakeOptions{AppData: true}
	result, err := parseAuthResult(valid, "tmp-1", options)
	if err != nil || result.AppID != "app-1" || result.SafeID != "safe-1" {
		t.Fatalf("valid result rejected: result=%#v err=%v", result, err)
	}

	for name, mutate := range map[string]func(map[string]any){
		"missing status":  func(value map[string]any) { delete(value, "status") },
		"wrong tmp":       func(value map[string]any) { value["tmp_id"] = "tmp-2" },
		"missing app":     func(value map[string]any) { value["app_id"] = "" },
		"missing safe":    func(value map[string]any) { value["safe_id"] = "" },
		"missing version": func(value map[string]any) { delete(value, "ver") },
		"missing ability": func(value map[string]any) {
			value["capabilities"] = []any{"sso"}
		},
		"extra ability": func(value map[string]any) {
			value["capabilities"] = []any{"sso", "app_data", "ap"}
		},
	} {
		t.Run(name, func(t *testing.T) {
			candidate := cloneMap(valid)
			mutate(candidate)
			if _, err := parseAuthResult(candidate, "tmp-1", options); err == nil {
				t.Fatal("invalid result accepted")
			}
		})
	}
}

func TestParseAuthResultAcceptsDeferredAppData(t *testing.T) {
	valid := map[string]any{
		"status":        "done",
		"tmp_id":        "tmp-1",
		"app_id":        "app-1",
		"safe_id":       "safe-1",
		"capabilities":  []any{"sso", "app_data"},
		"data_filekey":  []byte("opaque-filekey"),
		"data_deferred": true,
		"ver":           "7",
	}
	result, err := parseAuthResult(valid, "tmp-1", HandshakeOptions{AppData: true})
	if err != nil || !result.DataDeferred {
		t.Fatalf("valid deferred result rejected: result=%#v err=%v", result, err)
	}

	for name, mutate := range map[string]func(map[string]any){
		"missing filekey": func(value map[string]any) { delete(value, "data_filekey") },
		"inline content":  func(value map[string]any) { value["data_encrypted"] = []byte{1} },
		"false flag":      func(value map[string]any) { value["data_deferred"] = false },
		"wrong flag type": func(value map[string]any) { value["data_deferred"] = "true" },
	} {
		t.Run(name, func(t *testing.T) {
			candidate := cloneMap(valid)
			mutate(candidate)
			if _, err := parseAuthResult(candidate, "tmp-1", HandshakeOptions{AppData: true}); err == nil {
				t.Fatal("invalid deferred result accepted")
			}
		})
	}
}

func TestParseAuthResultAcceptsTrimmedLegacyAppData(t *testing.T) {
	valid := map[string]any{
		"status":       "done",
		"tmp_id":       "tmp-1",
		"app_id":       "app-1",
		"safe_id":      "safe-1",
		"capabilities": []any{"sso", "app_data"},
		"data_plain":   []byte("namespace metadata"),
		"ver":          "7",
	}
	result, err := parseAuthResult(valid, "tmp-1", HandshakeOptions{AppData: true})
	if err != nil ||
		result.DataDeferred ||
		len(result.DataFileKey) != 0 ||
		len(result.DataEncrypted) != 0 ||
		!bytes.Equal(result.DataPlain, []byte("namespace metadata")) {
		t.Fatalf("trimmed legacy result rejected: result=%#v err=%v", result, err)
	}
}

func TestParseAuthResultChecksTmpBeforeTerminalErr(t *testing.T) {
	terminal := map[string]any{
		"status": "error",
		"tmp_id": "tmp-1",
		"code":   AppAuthResultTooLargeCode,
	}
	_, err := parseAuthResult(terminal, "tmp-1", HandshakeOptions{})
	var apiErr *APIError
	if !errors.As(err, &apiErr) || apiErr.Code != AppAuthResultTooLargeCode {
		t.Fatalf("terminal result did not return APIError: %#v", err)
	}

	_, err = parseAuthResult(terminal, "tmp-2", HandshakeOptions{})
	if err == nil || errors.As(err, &apiErr) {
		t.Fatalf("mismatched terminal result was interpreted before tmp_id validation: %#v", err)
	}

	delete(terminal, "code")
	if _, err := parseAuthResult(terminal, "tmp-1", HandshakeOptions{}); err == nil {
		t.Fatal("terminal result without an error code was accepted")
	}
}

func TestRespErrPreservesReqIP(t *testing.T) {
	err := respErr(map[string]any{
		"_res":       "err",
		"code":       "ip_not_allowed",
		"request_ip": "2001:db8::10",
	})
	var apiErr *APIError
	if !errors.As(err, &apiErr) {
		t.Fatalf("expected APIError, got %T", err)
	}
	if apiErr.Code != "ip_not_allowed" ||
		apiErr.RequestIP != "2001:db8::10" ||
		apiErr.Error() != "ip_not_allowed (request_ip=2001:db8::10)" {
		t.Fatalf("unexpected API error: %#v (%v)", apiErr, apiErr)
	}
}

func TestParseAuthResultValidatesAPGrant(t *testing.T) {
	valid := map[string]any{
		"status":             "done",
		"tmp_id":             "tmp-1",
		"app_id":             "app-1",
		"safe_id":            "safe-1",
		"capabilities":       []any{"sso", "ap"},
		"user_enc_pk":        make([]byte, 32),
		"account_key_for_ap": make([]byte, 72),
		"ap_access": map[string]any{
			"full_access": false,
			"scopes":      []any{"2fa"},
		},
	}
	options := HandshakeOptions{AccessPoint: true, RequestedScopes: []string{"2fa"}}
	if _, err := parseAuthResult(valid, "tmp-1", options); err != nil {
		t.Fatalf("valid restricted AP result rejected: %v", err)
	}

	full := cloneMap(valid)
	full["ap_access"] = map[string]any{"full_access": true, "scopes": []any{}}
	if _, err := parseAuthResult(full, "tmp-1", HandshakeOptions{AccessPoint: true}); err != nil {
		t.Fatalf("valid full AP result rejected: %v", err)
	}

	wrongScope := cloneMap(valid)
	wrongScope["ap_access"] = map[string]any{
		"full_access": false,
		"scopes":      []any{"password"},
	}
	if _, err := parseAuthResult(wrongScope, "tmp-1", options); err == nil {
		t.Fatal("mismatched AP scope accepted")
	}
}

func TestWaitForAuthSignalsReadyBeforeReading(t *testing.T) {
	ready := false
	conn := &fakeAuthSocket{
		read: func() (int, []byte, error) {
			if !ready {
				t.Fatal("socket read started before ready callback")
			}
			return websocket.TextMessage, nil, errors.New("stop")
		},
	}
	_, err := (&Client{}).waitForAuthConn(
		conn,
		"tmp-1",
		HandshakeOptions{},
		time.Second,
		func() { ready = true },
	)
	if err == nil || !ready {
		t.Fatalf("unexpected wait result: ready=%v err=%v", ready, err)
	}
}

func TestRestoreRawFields(t *testing.T) {
	body := map[string]any{
		"data_filekey__": int8(1),
		"nested": map[string]any{
			"data_enc__": int64(0),
		},
	}
	if err := restoreRawFields(body, []any{[]byte("cipher"), []byte("filekey")}); err != nil {
		t.Fatalf("restore raw: %v", err)
	}
	if string(toBytes(body["data_filekey"])) != "filekey" {
		t.Fatal("top-level raw value not restored")
	}
	nested := body["nested"].(map[string]any)
	if string(toBytes(nested["data_enc"])) != "cipher" {
		t.Fatal("nested raw value not restored")
	}
}

func TestSealKPSProducesServerVerifiableFrame(t *testing.T) {
	client := NewClient()
	serverPk, serverSk, err := box.GenerateKey(rand.Reader)
	if err != nil {
		t.Fatal(err)
	}
	client.srvEncPk = serverPk
	frame, err := client.SealKPS(
		"app.42",
		"app_get_data",
		map[string]any{"request": "value"},
	)
	if err != nil {
		t.Fatal(err)
	}
	root, err := unpackMap(frame)
	if err != nil {
		t.Fatal(err)
	}
	kps := root["kps"].(map[string]any)
	openedHash, ok := sign.Open(nil, toBytes(root["seal"]), client.signPk)
	if !ok {
		t.Fatal("client KPS seal did not verify")
	}
	hash := sha256.Sum256(packCanonical(kps))
	if !bytes.Equal(openedHash, hash[:]) {
		t.Fatal("client KPS seal hash does not match canonical kps")
	}
	nonceBytes := toBytes(kps["n"])
	var nonce [24]byte
	copy(nonce[:], nonceBytes)
	plaintext, ok := box.Open(
		nil,
		toBytes(kps["box"]),
		&nonce,
		client.encPk,
		serverSk,
	)
	if !ok {
		t.Fatal("server could not open client KPS box")
	}
	inner, err := unpackMap(plaintext)
	if err != nil {
		t.Fatal(err)
	}
	header := inner["header"].(map[string]any)
	if asString(kps["id"]) != "app.42" ||
		asString(header["to"]) != "app_get_data" {
		t.Fatalf("unexpected outbound KPS content: kps=%v header=%v", kps, header)
	}
}

func TestOpenKPSVerifiesServerFrameAndRestoresRaw(t *testing.T) {
	client := NewClient()
	serverEncPk, serverEncSk, err := box.GenerateKey(rand.Reader)
	if err != nil {
		t.Fatal(err)
	}
	serverSignPk, serverSignSk, err := sign.GenerateKey(rand.Reader)
	if err != nil {
		t.Fatal(err)
	}
	client.srvEncPk = serverEncPk
	client.srvSignPk = serverSignPk

	nonce := [24]byte{}
	if _, err := rand.Read(nonce[:]); err != nil {
		t.Fatal(err)
	}
	inner := pack(map[string]any{
		"header": map[string]any{
			"ts":   time.Now().Unix(),
			"from": "app_filekey_result",
		},
		"body": map[string]any{
			"_res":           "ok",
			"status":         "done",
			"data_filekey__": int8(0),
		},
	})
	kps := map[string]any{
		"box": box.Seal(nil, inner, &nonce, client.encPk, serverEncSk),
		"n":   nonce[:],
		"raw": []any{[]byte("filekey")},
	}
	hash := sha256.Sum256(packCanonical(kps))
	frame := pack(map[string]any{
		"kps":  kps,
		"seal": sign.Sign(nil, hash[:], serverSignSk),
	})

	action, body, err := client.OpenKPS(frame)
	if err != nil {
		t.Fatal(err)
	}
	if action != "app_filekey_result" ||
		string(toBytes(body["data_filekey"])) != "filekey" {
		t.Fatalf("unexpected opened KPS: action=%q body=%v", action, body)
	}
}

func TestParseAuthorizationCallbackValidatesSingleResultAndTmp(t *testing.T) {
	client := NewClient()
	serverEncPk, serverEncSk, err := box.GenerateKey(rand.Reader)
	if err != nil {
		t.Fatal(err)
	}
	serverSignPk, serverSignSk, err := sign.GenerateKey(rand.Reader)
	if err != nil {
		t.Fatal(err)
	}
	client.srvEncPk = serverEncPk
	client.srvSignPk = serverSignPk

	frame := makeServerFrame(t, client, serverEncSk, serverSignSk, map[string]any{
		"_res":           "ok",
		"status":         "done",
		"tmp_id":         "tmp-1",
		"app_id":         "app-1",
		"safe_id":        "safe-1",
		"capabilities":   []any{"sso", "app_data"},
		"data_filekey":   []byte("opaque-filekey"),
		"data_encrypted": []byte("ciphertext"),
		"ver":            "7",
	})
	encoded := base64.RawURLEncoding.EncodeToString(frame)
	callback := "demo://auth-complete?state=abc&keylockr_result=" + url.QueryEscape(encoded)

	result, err := client.ParseAuthorizationCallback(
		callback,
		"tmp-1",
		HandshakeOptions{AppData: true},
	)
	if err != nil || result.TmpID != "tmp-1" || result.Version != "7" ||
		!bytes.Equal(result.DataFileKey, []byte("opaque-filekey")) {
		t.Fatalf("valid callback rejected: result=%#v err=%v", result, err)
	}
	if _, err := client.ParseAuthorizationCallback(
		callback+"&keylockr_result="+encoded,
		"tmp-1",
		HandshakeOptions{AppData: true},
	); err == nil {
		t.Fatal("duplicate keylockr_result accepted")
	}
	if _, err := client.ParseAuthorizationCallback(
		callback,
		"tmp-2",
		HandshakeOptions{AppData: true},
	); err == nil {
		t.Fatal("mismatched callback tmp_id accepted")
	}
	terminalFrame := makeServerFrame(t, client, serverEncSk, serverSignSk, map[string]any{
		"_res":   "ok",
		"status": "error",
		"tmp_id": "tmp-1",
		"code":   AppAuthResultTooLargeCode,
	})
	terminalCallback := "demo://auth-complete?keylockr_result=" +
		base64.RawURLEncoding.EncodeToString(terminalFrame)
	_, err = client.ParseAuthorizationCallback(
		terminalCallback,
		"tmp-1",
		HandshakeOptions{},
	)
	var terminalErr *APIError
	if !errors.As(err, &terminalErr) || terminalErr.Code != AppAuthResultTooLargeCode {
		t.Fatalf("terminal callback did not return the server error: %#v", err)
	}
	_, err = client.ParseAuthorizationCallback(
		terminalCallback,
		"tmp-2",
		HandshakeOptions{},
	)
	if err == nil || errors.As(err, &terminalErr) {
		t.Fatalf("terminal callback bypassed tmp_id validation: %#v", err)
	}
	stale := makeServerFrameAt(
		t,
		client,
		serverEncSk,
		serverSignSk,
		time.Now().Add(-maxServerAge-time.Second).Unix(),
		map[string]any{
			"_res":         "ok",
			"status":       "done",
			"tmp_id":       "tmp-1",
			"app_id":       "app-1",
			"safe_id":      "safe-1",
			"capabilities": []any{"sso"},
		},
	)
	if _, err := client.ParseAuthorizationCallback(
		"demo://auth-complete?keylockr_result="+base64.RawURLEncoding.EncodeToString(stale),
		"tmp-1",
		HandshakeOptions{},
	); err == nil {
		t.Fatal("stale callback KPS accepted")
	}
}

func makeServerFrame(
	t *testing.T,
	client *Client,
	serverEncSk *[32]byte,
	serverSignSk *[64]byte,
	body map[string]any,
) []byte {
	return makeServerFrameAt(
		t, client, serverEncSk, serverSignSk, time.Now().Unix(), body,
	)
}

func makeServerFrameAt(
	t *testing.T,
	client *Client,
	serverEncSk *[32]byte,
	serverSignSk *[64]byte,
	timestamp int64,
	body map[string]any,
) []byte {
	t.Helper()
	var nonce [24]byte
	if _, err := rand.Read(nonce[:]); err != nil {
		t.Fatal(err)
	}
	inner := pack(map[string]any{
		"header": map[string]any{"ts": timestamp, "from": "app_auth_result"},
		"body":   body,
	})
	kps := map[string]any{
		"box": box.Seal(nil, inner, &nonce, client.encPk, serverEncSk),
		"n":   nonce[:],
	}
	hash := sha256.Sum256(packCanonical(kps))
	return pack(map[string]any{
		"kps":  kps,
		"seal": sign.Sign(nil, hash[:], serverSignSk),
	})
}

func TestValidateServerTimestampBoundaries(t *testing.T) {
	now := time.Unix(2_000_000_000, 0)
	for _, timestamp := range []int64{
		now.Unix(),
		now.Add(-maxServerAge).Unix(),
		now.Add(maxServerLead).Unix(),
	} {
		if err := validateServerTimestamp(timestamp, now); err != nil {
			t.Fatalf("valid timestamp %d rejected: %v", timestamp, err)
		}
	}
	for _, timestamp := range []int64{
		now.Add(-maxServerAge - time.Second).Unix(),
		now.Add(maxServerLead + time.Second).Unix(),
	} {
		if err := validateServerTimestamp(timestamp, now); err == nil {
			t.Fatalf("invalid timestamp %d accepted", timestamp)
		}
	}
	if err := validateServerTimestamp("not-a-timestamp", now); err == nil {
		t.Fatal("non-numeric timestamp accepted")
	}
}

func TestBuildAuthorizationURI(t *testing.T) {
	direct, err := buildAuthorizationURI("tmp id", "demo://auth-complete")
	if err != nil {
		t.Fatal(err)
	}
	if direct != "keylockr://sso?return_url=demo%3A%2F%2Fauth-complete&tmp_id=tmp+id" {
		t.Fatalf("unexpected direct URI: %s", direct)
	}
	if _, err := buildAuthorizationURI("tmp", "file:///tmp/result"); err == nil {
		t.Fatal("blocked callback scheme accepted")
	}
	// BEGIN GENERATED SSO RETURN URL POLICY
	// return-url-blocked: about,blob,content,data,facetime,facetime-audio,file,intent,itms*,javascript,keylockr,mailto,market,sms,sso,tel,telprompt
	blockedSchemes := []string{
		"about", "blob", "content", "data", "facetime", "facetime-audio",
		"file", "intent", "itms", "itms-apps", "itms-services", "javascript",
		"keylockr", "mailto", "market", "sms", "sso", "tel",
		"telprompt",
	}
	blockedSchemePrefixes := []string{"itms"}
	for _, scheme := range blockedSchemes {
		for _, candidate := range []string{scheme, strings.ToUpper(scheme)} {
			callback := candidate + ":blocked"
			if _, err := buildAuthorizationURI("tmp", callback); err == nil {
				t.Fatalf("blocked callback scheme accepted: %s", callback)
			}
		}
	}
	for _, prefix := range blockedSchemePrefixes {
		for _, candidate := range []string{prefix + "-unlisted", strings.ToUpper(prefix) + "-UNLISTED"} {
			callback := candidate + ":blocked"
			if _, err := buildAuthorizationURI("tmp", callback); err == nil {
				t.Fatalf("blocked callback prefix accepted: %s", callback)
			}
		}
	}
	// END GENERATED SSO RETURN URL POLICY
	if _, err := buildAuthorizationURI("tmp", "/relative"); err == nil {
		t.Fatal("relative callback accepted")
	}
}

func cloneMap(source map[string]any) map[string]any {
	out := make(map[string]any, len(source))
	for key, value := range source {
		out[key] = value
	}
	return out
}

type fakeAuthSocket struct {
	read func() (int, []byte, error)
}

func (f *fakeAuthSocket) Close() error { return nil }

func (f *fakeAuthSocket) ReadMessage() (int, []byte, error) {
	return f.read()
}

func (f *fakeAuthSocket) SetReadDeadline(time.Time) error { return nil }
