/*
Copyright (c) 2026 Huawei Technologies Co., Ltd.
openFuyao is licensed under Mulan PSL v2.
You can use this software according to the terms and conditions of the Mulan PSL v2.
You may obtain a copy of Mulan PSL v2 at:
         http://license.coscl.org.cn/MulanPSL2
THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND,
EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT,
MERCHANTABILITY OR FIT FOR A PARTICULAR PURPOSE.
See the Mulan PSL v2 for more details.
*/

package remote

import (
	"crypto/rand"
	"crypto/rsa"
	"crypto/x509"
	"crypto/x509/pkix"
	"encoding/pem"
	"math/big"
	"os"
	"path/filepath"
	"testing"
	"time"

	"github.com/stretchr/testify/assert"
	"github.com/stretchr/testify/require"
)

func generateTestCA() ([]byte, *rsa.PrivateKey) {
	caKey, err := rsa.GenerateKey(rand.Reader, 2048)
	if err != nil {
		panic(err)
	}
	caTemplate := &x509.Certificate{
		SerialNumber:          big.NewInt(1),
		Subject:               pkix.Name{CommonName: "test-ca"},
		NotBefore:             time.Now(),
		NotAfter:              time.Now().Add(24 * time.Hour),
		IsCA:                  true,
		BasicConstraintsValid: true,
		KeyUsage:              x509.KeyUsageCertSign | x509.KeyUsageDigitalSignature,
	}
	caCertDER, err := x509.CreateCertificate(rand.Reader, caTemplate, caTemplate, &caKey.PublicKey, caKey)
	if err != nil {
		panic(err)
	}
	caCertPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: caCertDER})
	return caCertPEM, caKey
}

func generateTestClientCert(caCert *x509.Certificate, caKey *rsa.PrivateKey) ([]byte, []byte) {
	clientKey, err := rsa.GenerateKey(rand.Reader, 2048)
	if err != nil {
		panic(err)
	}
	clientTemplate := &x509.Certificate{
		SerialNumber: big.NewInt(2),
		Subject:      pkix.Name{CommonName: "test-client"},
		NotBefore:    time.Now(),
		NotAfter:     time.Now().Add(24 * time.Hour),
		KeyUsage:     x509.KeyUsageDigitalSignature,
	}
	clientCertDER, err := x509.CreateCertificate(rand.Reader, clientTemplate, caCert, &clientKey.PublicKey, caKey)
	if err != nil {
		panic(err)
	}
	clientCertPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: clientCertDER})
	clientKeyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(clientKey)})
	return clientCertPEM, clientKeyPEM
}

func validTestConfig() *Config {
	caCertPEM, caKey := generateTestCA()
	caTemplate := &x509.Certificate{
		SerialNumber:          big.NewInt(1),
		Subject:               pkix.Name{CommonName: "test-ca"},
		NotBefore:             time.Now(),
		NotAfter:              time.Now().Add(24 * time.Hour),
		IsCA:                  true,
		BasicConstraintsValid: true,
		KeyUsage:              x509.KeyUsageCertSign | x509.KeyUsageDigitalSignature,
	}
	clientCertPEM, clientKeyPEM := generateTestClientCert(caTemplate, caKey)
	return &Config{
		ServerAddr: "https://test-server:8443",
		CACert:     caCertPEM,
		ClientCert: clientCertPEM,
		ClientKey:  clientKeyPEM,
	}
}

func TestNew(t *testing.T) {
	t.Run("成功场景", func(t *testing.T) {
		config := validTestConfig()
		mgr, err := New(config)
		require.NoError(t, err)
		require.NotNil(t, mgr)
	})

	t.Run("客户端证书与密钥不匹配", func(t *testing.T) {
		otherKey, err := rsa.GenerateKey(rand.Reader, 2048)
		require.NoError(t, err)
		otherKeyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(otherKey)})
		config := validTestConfig()
		config.ClientKey = otherKeyPEM
		_, err = New(config)
		require.Error(t, err)
		assert.Contains(t, err.Error(), "load client cert/key")
	})

	t.Run("CA证书无效", func(t *testing.T) {
		config := validTestConfig()
		config.CACert = []byte("not-a-valid-pem")
		_, err := New(config)
		require.Error(t, err)
		assert.Contains(t, err.Error(), "failed to append CA cert")
	})

	t.Run("空客户端证书和密钥", func(t *testing.T) {
		config := validTestConfig()
		config.ClientCert = nil
		config.ClientKey = nil
		_, err := New(config)
		require.Error(t, err)
		assert.Contains(t, err.Error(), "load client cert/key")
	})
}

func writeConfigFiles(t *testing.T, dir string, config *Config, skipFiles ...string) {
	t.Helper()
	skipSet := make(map[string]bool, len(skipFiles))
	for _, f := range skipFiles {
		skipSet[f] = true
	}
	if !skipSet["serverAddr"] {
		require.NoError(t, os.WriteFile(filepath.Join(dir, "serverAddr"), []byte("https://test-server:8443"), 0644))
	}
	if !skipSet["ca.crt"] {
		require.NoError(t, os.WriteFile(filepath.Join(dir, "ca.crt"), config.CACert, 0644))
	}
	if !skipSet["client.crt"] {
		require.NoError(t, os.WriteFile(filepath.Join(dir, "client.crt"), config.ClientCert, 0644))
	}
	if !skipSet["client.key"] {
		require.NoError(t, os.WriteFile(filepath.Join(dir, "client.key"), config.ClientKey, 0644))
	}
}

func TestLoadConfigFromDir成功场景(t *testing.T) {
	config := validTestConfig()
	dir := t.TempDir()
	writeConfigFiles(t, dir, config)

	loaded, err := LoadConfigFromDir(dir)
	require.NoError(t, err)
	require.NotNil(t, loaded)
	assert.Equal(t, "https://test-server:8443", loaded.ServerAddr)
	assert.Equal(t, config.CACert, loaded.CACert)
	assert.Equal(t, config.ClientCert, loaded.ClientCert)
	assert.Equal(t, config.ClientKey, loaded.ClientKey)
}

func TestLoadConfigFromDir文件缺失(t *testing.T) {
	missingFileCases := []struct {
		name     string
		skipFile string
		wantErr  string
	}{
		{"serverAddr文件缺失", "serverAddr", "read serverAddr"},
		{"ca.crt文件缺失", "ca.crt", "read ca.crt"},
		{"client.crt文件缺失", "client.crt", "read client.crt"},
		{"client.key文件缺失", "client.key", "read client.key"},
	}
	for _, tc := range missingFileCases {
		t.Run(tc.name, func(t *testing.T) {
			config := validTestConfig()
			dir := t.TempDir()
			writeConfigFiles(t, dir, config, tc.skipFile)

			_, err := LoadConfigFromDir(dir)
			require.Error(t, err)
			assert.Contains(t, err.Error(), tc.wantErr)
		})
	}
}

func TestLoadConfigFromDirTrimSpace(t *testing.T) {
	config := validTestConfig()
	dir := t.TempDir()
	writeConfigFiles(t, dir, config)
	require.NoError(t, os.WriteFile(filepath.Join(dir, "serverAddr"), []byte("  https://test-server:8443  \n"), 0644))

	loaded, err := LoadConfigFromDir(dir)
	require.NoError(t, err)
	assert.Equal(t, "https://test-server:8443", loaded.ServerAddr)
}

func TestNewFromDir(t *testing.T) {
	t.Run("成功场景", func(t *testing.T) {
		config := validTestConfig()
		dir := t.TempDir()
		writeConfigFiles(t, dir, config)

		mgr, err := NewFromDir(dir)
		require.NoError(t, err)
		require.NotNil(t, mgr)
	})

	t.Run("目录不存在", func(t *testing.T) {
		_, err := NewFromDir("/nonexistent/path")
		require.Error(t, err)
		assert.Contains(t, err.Error(), "load config from dir")
	})

	t.Run("目录存在但证书无效", func(t *testing.T) {
		dir := t.TempDir()
		require.NoError(t, os.WriteFile(filepath.Join(dir, "serverAddr"), []byte("https://test-server:8443"), 0644))
		require.NoError(t, os.WriteFile(filepath.Join(dir, "ca.crt"), []byte("valid-ca-but-will-be-overridden"), 0644))
		require.NoError(t, os.WriteFile(filepath.Join(dir, "client.crt"), []byte("not-a-cert"), 0644))
		require.NoError(t, os.WriteFile(filepath.Join(dir, "client.key"), []byte("not-a-key"), 0644))

		_, err := NewFromDir(dir)
		require.Error(t, err)
		assert.Contains(t, err.Error(), "load client cert/key")
	})
}