package utils
import (
"bytes"
"context"
"crypto/aes"
"crypto/cipher"
"crypto/hmac"
"crypto/sha256"
"encoding/hex"
"fmt"
"log"
"github.com/go-redis/redis/v8"
"golang.org/x/crypto/pbkdf2"
)
var ctx = context.Background()
const (
redisAddr = "110.41.7.101:6388"
redisPassword = "Test@Device@2024"
redisDB = 9
)
func initRedisClient() *redis.Client {
client := redis.NewClient(&redis.Options{
Addr: redisAddr,
Password: redisPassword,
DB: redisDB,
})
_, err := client.Ping(ctx).Result()
if err != nil {
log.Fatalf("Failed to connect to Redis: %v", err)
}
return client
}
func GetEncryptMode(deviceID string) (string, error) {
client := initRedisClient()
defer client.Close()
encryptKey := fmt.Sprintf("iot%s", deviceID)
encrypt, err := client.HGet(ctx, encryptKey, "encryption").Result()
if err != nil {
if err == redis.Nil {
return "plaintext", nil
}
return "plaintext", fmt.Errorf("failed to get encryption mode from Redis: %w", err)
}
return encrypt, nil
}
func getRemoteControlParams(client *redis.Client, deviceID string) ([]byte, []byte, []byte, error) {
sn1Key := fmt.Sprintf("iot%s", deviceID)
sn2Key := fmt.Sprintf("iot%s", deviceID)
pskKey := fmt.Sprintf("iot%s", deviceID)
sn1Hex, err := client.HGet(ctx, sn1Key, "sn1").Result()
if err != nil {
return nil, nil, nil, fmt.Errorf("failed to get SN1 from Redis: %v", err)
}
sn2Hex, err := client.HGet(ctx, sn2Key, "sn2").Result()
if err != nil {
return nil, nil, nil, fmt.Errorf("failed to get SN2 from Redis: %v", err)
}
pskHex, err := client.HGet(ctx, pskKey, "psk").Result()
if err != nil {
return nil, nil, nil, fmt.Errorf("failed to get PSK from Redis: %v", err)
}
sn1, err := hex.DecodeString(sn1Hex)
if err != nil {
return nil, nil, nil, fmt.Errorf("sn1 解码失败: %v", err)
}
sn2, err := hex.DecodeString(sn2Hex)
if err != nil {
return nil, nil, nil, fmt.Errorf("sn2 解码失败: %v", err)
}
psk, err := hex.DecodeString(pskHex)
if err != nil {
return nil, nil, nil, fmt.Errorf("PSK 解码失败: %v", err)
}
if len(sn1) != 8 || len(sn2) != 8 {
return nil, nil, nil, fmt.Errorf("SN1/SN2 must be 8 bytes after hex decoding")
}
return sn1, sn2, psk, nil
}
func deriveKey(password, salt []byte, iterations, keyLen int) []byte {
return pbkdf2.Key(password, salt, iterations, keyLen, sha256.New)
}
func pkcs5Padding(data []byte, blockSize int) []byte {
padding := blockSize - len(data)%blockSize
padtext := bytes.Repeat([]byte{byte(padding)}, padding)
return append(data, padtext...)
}
func pkcs5Unpadding(data []byte) ([]byte, error) {
if len(data) == 0 {
return nil, fmt.Errorf("data is empty")
}
paddingLen := int(data[len(data)-1])
if paddingLen > len(data) || paddingLen == 0 {
return nil, fmt.Errorf("invalid padding length")
}
for i := len(data) - paddingLen; i < len(data); i++ {
if data[i] != byte(paddingLen) {
return nil, fmt.Errorf("invalid padding bytes")
}
}
return data[:len(data)-paddingLen], nil
}
func encryptAES128CBC(plaintext, key, iv []byte) ([]byte, error) {
block, err := aes.NewCipher(key)
if err != nil {
return nil, err
}
paddedData := pkcs5Padding(plaintext, aes.BlockSize)
mode := cipher.NewCBCEncrypter(block, iv)
ciphertext := make([]byte, len(paddedData))
mode.CryptBlocks(ciphertext, paddedData)
return ciphertext, nil
}
func decryptAES128CBC(ciphertext, key, iv []byte) ([]byte, error) {
block, err := aes.NewCipher(key)
if err != nil {
return nil, err
}
mode := cipher.NewCBCDecrypter(block, iv)
decrypted := make([]byte, len(ciphertext))
mode.CryptBlocks(decrypted, ciphertext)
return pkcs5Unpadding(decrypted)
}
func calculateHMAC(data, secret []byte) []byte {
h := hmac.New(sha256.New, secret)
h.Write(data)
return h.Sum(nil)
}
func buildCoAPMessage(coapHeader, payload, hmac []byte) []byte {
return append(append(coapHeader, payload...), hmac...)
}
func EncryptBodyAES(devId string, payload []byte) []byte {
client := initRedisClient()
defer client.Close()
sn1, sn2, psk, err := getRemoteControlParams(client, devId)
if err != nil {
log.Fatalf("Failed to get encryption parameters: %v", err)
}
salt := append(sn1, sn2...)
digest := deriveKey(psk, salt, 1, 32)
aesKey := digest[:16]
iv := digest[16:32]
ciphertext, err := encryptAES128CBC(payload, aesKey, iv)
if err != nil {
log.Fatalf("Encryption error: %v", err)
}
return ciphertext
}
func EncryptDataAESMac(devId string, payload []byte, coapHeader []byte) []byte {
client := initRedisClient()
defer client.Close()
sn1, sn2, psk, err := getRemoteControlParams(client, devId)
if err != nil {
log.Fatalf("Failed to get encryption parameters: %v", err)
}
salt := append(sn1, sn2...)
digest := deriveKey(psk, salt, 1, 32)
aesKey := digest[:16]
iv := digest[16:32]
ciphertext, err := encryptAES128CBC(payload, aesKey, iv)
if err != nil {
log.Fatalf("Encryption error: %v", err)
}
hmacSecret := deriveKey(digest[:16], salt, 1, 32)
inputData := append(coapHeader, ciphertext...)
mac := calculateHMAC(inputData, hmacSecret)
finalMessage := buildCoAPMessage(coapHeader, ciphertext, mac)
return finalMessage
}
func CheckMac(devId string, payload []byte, coapHeader []byte) (bool, error) {
client := initRedisClient()
defer client.Close()
sn1, sn2, psk, err := getRemoteControlParams(client, devId)
if err != nil {
log.Fatalf("Failed to get encryption parameters: %v", err)
}
salt := append(sn1, sn2...)
digest := deriveKey(psk, salt, 1, 32)
hmacSecret := deriveKey(digest[:16], salt, 1, 32)
if len(payload) < 32 {
log.Fatalf("Decryption error: %v", err)
}
msg := len(payload) - 32
enMac := payload[msg:]
enpaylod := payload[:msg]
inputData := append(coapHeader, enpaylod...)
mac := calculateHMAC(inputData, hmacSecret)
if !hmac.Equal(mac, enMac) {
return false, fmt.Errorf("HMAC verification failed")
}
return true, nil
}
func DecryptData(devId string, payload []byte) ([]byte, error) {
client := initRedisClient()
defer client.Close()
sn1, sn2, psk, err := getRemoteControlParams(client, devId)
if err != nil {
log.Fatalf("Failed to get encryption parameters: %v", err)
}
salt := append(sn1, sn2...)
digest := deriveKey(psk, salt, 1, 32)
aesKey := digest[:16]
iv := digest[16:32]
if len(payload) < 32 {
log.Fatalf("Decryption error: %v", err)
}
decrypted, err := decryptAES128CBC(payload, aesKey, iv)
if err != nil {
log.Fatalf("Decryption error: %v", err)
}
fmt.Printf("Decrypted Payload: %s\n", decrypted)
return decrypted, err
}
func DecryptDataAESMac(devId string, payload []byte, coapHeader []byte) []byte {
client := initRedisClient()
defer client.Close()
sn1, sn2, psk, err := getRemoteControlParams(client, devId)
if err != nil {
log.Fatalf("Failed to get encryption parameters: %v", err)
}
salt := append(sn1, sn2...)
digest := deriveKey(psk, salt, 1, 32)
aesKey := digest[:16]
iv := digest[16:32]
hmacSecret := deriveKey(digest[:16], salt, 1, 32)
if len(payload) < 32 {
log.Fatalf("Decryption error: %v", err)
}
msg := len(payload) - 32
enMac := payload[msg:]
enpaylod := payload[:msg]
inputData := append(coapHeader, enpaylod...)
mac := calculateHMAC(inputData, hmacSecret)
if !hmac.Equal(mac, enMac) {
log.Fatalf("HMAC verification failed")
}
decrypted, err := decryptAES128CBC(enpaylod, aesKey, iv)
if err != nil {
log.Fatalf("Decryption error: %v", err)
}
fmt.Printf("Decrypted Payload: %s\n", decrypted)
return decrypted
}