package main

import (
	"bytes"
	"context"
	"crypto/ed25519"
	"crypto/sha256"
	"encoding/base64"
	"encoding/hex"
	"encoding/json"
	"errors"
	"os"
	"testing"
	"time"
)

const delegatedFixtureSigningSeedHex = "9d61b19deffd5a60ba844af492ec2cc4" +
	"4449c5697b326919703bac031cae7f60"

type delegatedFixture struct {
	AuthRequestHex             string `json:"auth_request_hex"`
	RequestDigestHex           string `json:"request_digest_hex"`
	BackendAssertionHex        string `json:"backend_assertion_hex"`
	SuccessWrapperHex          string `json:"success_wrapper_hex"`
	KLSignPublicKeyHex         string `json:"kl_sign_public_key_hex"`
	KLEncryptPublicKeyHex      string `json:"kl_encrypt_public_key_hex"`
	ServiceEncryptSecretKeyHex string `json:"service_encrypt_secret_key_hex"`
	FinishSecretHex            string `json:"finish_secret_hex"`
	FinishSecretHashHex        string `json:"finish_secret_hash_hex"`
	AssertionDigestHex         string `json:"assertion_digest_hex"`
	StartWallTime              int64  `json:"start_wall_time"`
	ChallengeExpiresAt         int64  `json:"challenge_expires_at"`
	RetentionExpiresAt         int64  `json:"retention_expires_at"`
	AssertionIssuedAt          int64  `json:"assertion_iat"`
	AssertionExpiresAt         int64  `json:"assertion_exp"`
	SafeID                     string `json:"safe_id"`
	AppID                      string `json:"app_id"`
	UserNickname               string `json:"user_nickname"`
}

func TestDelegatedStartMatchesCanonicalFixture(t *testing.T) {
	fixture := loadDelegatedFixture(t)
	challenge := make([]byte, delegatedRandomBytes)
	for index := range challenge {
		challenge[index] = byte(index)
	}
	finishSecret := decodeFixtureHex(t, fixture.FinishSecretHex)
	random := bytes.NewReader(append(challenge, finishSecret...))
	privateKey := ed25519.NewKeyFromSeed(decodeFixtureHex(t, delegatedFixtureSigningSeedHex))
	result, err := startDelegatedSSO(
		DelegatedStartOptions{
			AppTag: "cing", ReturnURL: "cing://sso/callback", Capabilities: []string{"sso"},
		},
		privateKey,
		time.Unix(fixture.StartWallTime, 0),
		random,
	)
	if err != nil {
		t.Fatal(err)
	}
	requestEncoded := result.DeepLink[len("keylockr://sso-request?request="):]
	authReq, err := base64.RawURLEncoding.DecodeString(requestEncoded)
	if err != nil {
		t.Fatal(err)
	}
	if got := hex.EncodeToString(authReq); got != fixture.AuthRequestHex {
		t.Fatalf("auth request fixture mismatch\ngot:  %s\nwant: %s", got, fixture.AuthRequestHex)
	}
	if result.RequestDigest != base64.RawURLEncoding.EncodeToString(decodeFixtureHex(t, fixture.RequestDigestHex)) {
		t.Fatal("request digest response changed")
	}
	if result.FinishSecret != base64.RawURLEncoding.EncodeToString(finishSecret) {
		t.Fatal("finish secret response changed")
	}
	if result.ExpiresAt != fixture.StartWallTime+600 || result.RetentionExpiresAt != fixture.StartWallTime+720 ||
		result.Pending.ChallengeExpiresAt != fixture.ChallengeExpiresAt ||
		result.Pending.RetentionExpiresAt != fixture.RetentionExpiresAt {
		t.Fatal("delegated deadlines must share one start wall-time snapshot")
	}
	if got := hex.EncodeToString(result.Pending.FinishSecretHash[:]); got != fixture.FinishSecretHashHex {
		t.Fatalf("finish secret hash got %s", got)
	}
}

func TestDelegatedStartRejectsCapabilityWithoutDeviceName(t *testing.T) {
	privateKey := ed25519.NewKeyFromSeed(decodeFixtureHex(t, delegatedFixtureSigningSeedHex))
	_, err := StartDelegatedSSO(DelegatedStartOptions{
		AppTag: "cing", ReturnURL: "cing://sso/callback", Capabilities: []string{"sso", "app_data"},
		ClientSignPK: bytes.Repeat([]byte{1}, delegatedPublicKeyBytes),
		ClientEncPK:  bytes.Repeat([]byte{2}, delegatedPublicKeyBytes),
	}, privateKey)
	if err == nil {
		t.Fatal("AppData request without a device name must fail before signing")
	}
}

func TestDelegatedStartRejectsDuplicateCapabilitiesAndScopes(t *testing.T) {
	privateKey := ed25519.NewKeyFromSeed(decodeFixtureHex(t, delegatedFixtureSigningSeedHex))
	deviceName := "Fixture device"
	base := DelegatedStartOptions{
		AppTag: "cing", ReturnURL: "cing://sso/callback", DeviceName: &deviceName,
		ClientSignPK: bytes.Repeat([]byte{1}, delegatedPublicKeyBytes),
		ClientEncPK:  bytes.Repeat([]byte{2}, delegatedPublicKeyBytes),
	}
	base.Capabilities = []string{"sso", "app_data", "app_data"}
	if _, err := StartDelegatedSSO(base, privateKey); err == nil {
		t.Fatal("duplicate capability must fail before signing")
	}
	base.Capabilities = []string{"sso", "ap"}
	base.RequestedScopes = []string{"2fa", "2fa"}
	if _, err := StartDelegatedSSO(base, privateKey); err == nil {
		t.Fatal("duplicate AP scope must fail before signing")
	}
}

func TestDelegatedCallbackAndFinishFixture(t *testing.T) {
	fixture := loadDelegatedFixture(t)
	pending := fixturePending(t, fixture)
	wrapper := decodeFixtureHex(t, fixture.SuccessWrapperHex)
	callback, err := ParseDelegatedCallback(
		base64.RawURLEncoding.EncodeToString(wrapper),
		pending.RequestDigest,
	)
	if err != nil {
		t.Fatal(err)
	}
	if callback.Cancelled || !bytes.Equal(callback.BackendAssertion, decodeFixtureHex(t, fixture.BackendAssertionHex)) {
		t.Fatal("canonical callback was not parsed")
	}
	verified, err := VerifyDelegatedFinish(
		pending,
		base64.RawURLEncoding.EncodeToString(decodeFixtureHex(t, fixture.FinishSecretHex)),
		callback.BackendAssertion,
		decodeFixtureHex(t, fixture.KLSignPublicKeyHex),
		decodeFixtureHex(t, fixture.KLEncryptPublicKeyHex),
		decodeFixtureHex(t, fixture.ServiceEncryptSecretKeyHex),
		time.Unix(fixture.AssertionIssuedAt+70, 0),
	)
	if err != nil {
		t.Fatal(err)
	}
	if verified.Identity.SafeID != fixture.SafeID || verified.Identity.AppID != fixture.AppID ||
		verified.Identity.Nickname != fixture.UserNickname ||
		verified.Identity.IssuedAt != fixture.AssertionIssuedAt ||
		verified.Identity.ExpiresAt != fixture.AssertionExpiresAt {
		t.Fatalf("unexpected verified identity: %#v", verified.Identity)
	}
	if got := hex.EncodeToString(verified.AssertionDigest[:]); got != fixture.AssertionDigestHex {
		t.Fatalf("assertion digest got %s", got)
	}
}

func TestDelegatedFinishRejectsInterceptionAndBoundary(t *testing.T) {
	fixture := loadDelegatedFixture(t)
	pending := fixturePending(t, fixture)
	pending.RetentionExpiresAt = fixture.AssertionExpiresAt + 61
	assertion := decodeFixtureHex(t, fixture.BackendAssertionHex)
	verify := func(secret string, now int64) error {
		_, err := VerifyDelegatedFinish(
			pending, secret, assertion,
			decodeFixtureHex(t, fixture.KLSignPublicKeyHex),
			decodeFixtureHex(t, fixture.KLEncryptPublicKeyHex),
			decodeFixtureHex(t, fixture.ServiceEncryptSecretKeyHex),
			time.Unix(now, 0),
		)
		return err
	}
	wrongSecret := base64.RawURLEncoding.EncodeToString(bytes.Repeat([]byte{0x42}, delegatedRandomBytes))
	if !errors.Is(verify(wrongSecret, fixture.AssertionExpiresAt+59), ErrDelegatedAuthentication) {
		t.Fatal("interceptor without finish secret must be rejected")
	}
	secret := base64.RawURLEncoding.EncodeToString(decodeFixtureHex(t, fixture.FinishSecretHex))
	if !errors.Is(verify(secret+"=", fixture.AssertionExpiresAt+59), ErrDelegatedAuthentication) {
		t.Fatal("padded finish secret must be rejected")
	}
	if !errors.Is(verify(secret, fixture.AssertionExpiresAt+60), ErrDelegatedAuthentication) {
		t.Fatal("assertion exp + 60 must be rejected")
	}
	tampered := append([]byte(nil), assertion...)
	tampered[len(tampered)-1] ^= 1
	if _, err := VerifyDelegatedFinish(
		pending, secret, tampered,
		decodeFixtureHex(t, fixture.KLSignPublicKeyHex),
		decodeFixtureHex(t, fixture.KLEncryptPublicKeyHex),
		decodeFixtureHex(t, fixture.ServiceEncryptSecretKeyHex),
		time.Unix(fixture.AssertionIssuedAt+70, 0),
	); !errors.Is(err, ErrDelegatedAuthentication) {
		t.Fatal("tampered assertion must return generic failure")
	}
}

func TestDelegatedCallbackRejectsDigestMismatchAndOversize(t *testing.T) {
	fixture := loadDelegatedFixture(t)
	pending := fixturePending(t, fixture)
	encoded := base64.RawURLEncoding.EncodeToString(decodeFixtureHex(t, fixture.SuccessWrapperHex))
	wrongDigest := pending.RequestDigest
	wrongDigest[0] ^= 1
	if _, err := ParseDelegatedCallback(encoded, wrongDigest); !errors.Is(err, ErrDelegatedAuthentication) {
		t.Fatal("callback digest mismatch must fail before finish")
	}
	if _, err := ParseDelegatedCallback(stringsOfSize(delegatedMaxEncodedResultChars+1), pending.RequestDigest); !errors.Is(err, ErrDelegatedAuthentication) {
		t.Fatal("oversized callback must fail")
	}
	if _, err := ParseDelegatedCallback(encoded+"=", pending.RequestDigest); !errors.Is(err, ErrDelegatedAuthentication) {
		t.Fatal("padded callback must fail canonical base64url validation")
	}
	cancelled, err := ParseDelegatedCallback("cancel", pending.RequestDigest)
	if err != nil || !cancelled.Cancelled {
		t.Fatal("fixed cancel callback must be accepted without assertion")
	}
}

func TestDelegatedCallbackURLRejectsMalformedOrAmbiguousQuery(t *testing.T) {
	fixture := loadDelegatedFixture(t)
	pending := fixturePending(t, fixture)
	encoded := base64.RawURLEncoding.EncodeToString(decodeFixtureHex(t, fixture.SuccessWrapperHex))
	valid, err := ParseDelegatedCallbackURL("cing://login?keylockr_result="+encoded, pending.RequestDigest)
	if err != nil || valid.Cancelled || len(valid.BackendAssertion) == 0 {
		t.Fatal("valid callback URL was rejected")
	}
	for _, invalid := range []string{
		"cing://login?keylockr_result=%ZZ",
		"cing://login?keylockr_result=" + encoded + "&keylockr_result=" + encoded,
		"cing://login?keylockr_result=" + encoded + "&extra=1",
	} {
		if _, err := ParseDelegatedCallbackURL(invalid, pending.RequestDigest); !errors.Is(err, ErrDelegatedAuthentication) {
			t.Fatalf("ambiguous callback URL was accepted: %q", invalid)
		}
	}
}

func TestCompleteDelegatedSSOUsesStableIdempotencyInput(t *testing.T) {
	fixture := loadDelegatedFixture(t)
	pending := fixturePending(t, fixture)
	store := &recordingCompletionStore{}
	result, err := CompleteDelegatedSSO(
		context.Background(), store, pending,
		base64.RawURLEncoding.EncodeToString(decodeFixtureHex(t, fixture.FinishSecretHex)),
		decodeFixtureHex(t, fixture.BackendAssertionHex),
		decodeFixtureHex(t, fixture.KLSignPublicKeyHex),
		decodeFixtureHex(t, fixture.KLEncryptPublicKeyHex),
		decodeFixtureHex(t, fixture.ServiceEncryptSecretKeyHex),
		time.Unix(fixture.AssertionIssuedAt+70, 0),
	)
	if err != nil {
		t.Fatal(err)
	}
	if result.SessionID != "session-fixture" || store.request.Identity.SafeID != fixture.SafeID ||
		store.request.RetentionExpiresAt != fixture.RetentionExpiresAt {
		t.Fatalf("unexpected completion request: %#v", store.request)
	}
	store.returnNickname = "wrong nickname"
	if _, err := CompleteDelegatedSSO(
		context.Background(), store, pending,
		base64.RawURLEncoding.EncodeToString(decodeFixtureHex(t, fixture.FinishSecretHex)),
		decodeFixtureHex(t, fixture.BackendAssertionHex),
		decodeFixtureHex(t, fixture.KLSignPublicKeyHex),
		decodeFixtureHex(t, fixture.KLEncryptPublicKeyHex),
		decodeFixtureHex(t, fixture.ServiceEncryptSecretKeyHex),
		time.Unix(fixture.AssertionIssuedAt+70, 0),
	); !errors.Is(err, ErrDelegatedAuthentication) {
		t.Fatal("completion store identity mismatch must fail closed")
	}
}

type recordingCompletionStore struct {
	request        DelegatedCompletionRequest
	returnNickname string
}

func (s *recordingCompletionStore) CompleteDelegatedSSO(
	_ context.Context,
	request DelegatedCompletionRequest,
) (DelegatedSessionResult, error) {
	s.request = request
	nickname := request.Identity.Nickname
	if s.returnNickname != "" {
		nickname = s.returnNickname
	}
	return DelegatedSessionResult{
		SessionID: "session-fixture", SafeID: request.Identity.SafeID, Nickname: nickname,
	}, nil
}

func loadDelegatedFixture(t *testing.T) delegatedFixture {
	t.Helper()
	raw, err := os.ReadFile("testdata/delegated_sso_v1.json")
	if err != nil {
		t.Fatal(err)
	}
	var fixture delegatedFixture
	if err := json.Unmarshal(raw, &fixture); err != nil {
		t.Fatal(err)
	}
	return fixture
}

func fixturePending(t *testing.T, fixture delegatedFixture) DelegatedPending {
	t.Helper()
	challenge := make([]byte, delegatedRandomBytes)
	for index := range challenge {
		challenge[index] = byte(index)
	}
	return DelegatedPending{
		AppTag: "cing", ChallengeHash: sha256.Sum256(challenge),
		FinishSecretHash: sha256.Sum256(decodeFixtureHex(t, fixture.FinishSecretHex)),
		RequestDigest:    array32(t, decodeFixtureHex(t, fixture.RequestDigestHex)),
		Capabilities:     []string{"sso"}, ChallengeExpiresAt: fixture.ChallengeExpiresAt,
		RetentionExpiresAt: fixture.RetentionExpiresAt,
	}
}

func array32(t *testing.T, raw []byte) [sha256.Size]byte {
	t.Helper()
	if len(raw) != sha256.Size {
		t.Fatalf("expected %d bytes, got %d", sha256.Size, len(raw))
	}
	var value [sha256.Size]byte
	copy(value[:], raw)
	return value
}

func decodeFixtureHex(t *testing.T, value string) []byte {
	t.Helper()
	decoded, err := hex.DecodeString(value)
	if err != nil {
		t.Fatal(err)
	}
	return decoded
}

func stringsOfSize(size int) string {
	return string(bytes.Repeat([]byte{'a'}, size))
}
