Files
heicode-mananger/heicode/middleware/device_signature_test.go
T

352 lines
11 KiB
Go

package middleware
import (
"bytes"
"crypto/ed25519"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/heicode/manager/model"
"github.com/heicode/manager/setting/operation_setting"
)
// rfc8032TestKeypair is the Ed25519 test vector 1 from RFC 8032. PUBLIC,
// only for reproducible test cases. Never use as a real signing key.
const rfc8032SeedHex = "9d61b19deffd5a60ba844af492ec2cc44449c5697b326919703bac031cae7f60"
func loadRFC8032PrivateKey(t *testing.T) ed25519.PrivateKey {
t.Helper()
seed, err := hex.DecodeString(rfc8032SeedHex)
if err != nil {
t.Fatalf("decode seed: %v", err)
}
return ed25519.NewKeyFromSeed(seed)
}
// signedReq builds a fully-formed *gin.Context whose request carries
// matching X-Heicode-* headers + Authorization Bearer for the given
// token. Callers can mutate fields before passing to
// VerifyDeviceSignatureIfRequired to test tampering scenarios.
type signedReq struct {
Method string
Path string
Body string
Timestamp int64
Nonce string
Fingerprint string
}
func defaultSignedReq() signedReq {
return signedReq{
Method: "POST",
Path: "/v1/messages",
Body: `{"model":"claude-sonnet-4-6","max_tokens":10}`,
Timestamp: time.Now().UnixMilli(),
Nonce: randomNonce(),
Fingerprint: strings.Repeat("a", 64),
}
}
func randomNonce() string {
// 16-byte nonce as 32-char lowercase hex. Tests don't need crypto-
// strong randomness, but each invocation MUST differ so replay
// detection asserts don't false-fire. Use crypto/rand for guaranteed
// uniqueness across rapid same-nanosecond calls.
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
// fall back to nanos-derived if rand is somehow broken
now := time.Now().UnixNano()
for i := 0; i < 16; i++ {
b[i] = byte(now >> uint(i*4))
}
}
return hex.EncodeToString(b[:])
}
func canonicalString(r signedReq, bodyHash string) string {
return strings.Join([]string{
strings.ToUpper(r.Method),
r.Path,
fmt.Sprint(r.Timestamp),
r.Nonce,
r.Fingerprint,
bodyHash,
}, "\n")
}
func makeSignedContext(t *testing.T, r signedReq, priv ed25519.PrivateKey) *gin.Context {
t.Helper()
bodySum := sha256.Sum256([]byte(r.Body))
bodyHash := hex.EncodeToString(bodySum[:])
canonical := canonicalString(r, bodyHash)
digest := sha256.Sum256([]byte(canonical))
sig := ed25519.Sign(priv, digest[:])
req := httptest.NewRequest(r.Method, r.Path, bytes.NewBufferString(r.Body))
req.Header.Set("Authorization", "Bearer sk-test")
req.Header.Set(HeaderDeviceID, "device-test-1")
req.Header.Set(HeaderTimestamp, fmt.Sprint(r.Timestamp))
req.Header.Set(HeaderNonce, r.Nonce)
req.Header.Set(HeaderFingerprint, r.Fingerprint)
req.Header.Set(HeaderSignature, base64.StdEncoding.EncodeToString(sig))
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
return c
}
func makeBoundToken(priv ed25519.PrivateKey) *model.Token {
pub := priv.Public().(ed25519.PublicKey)
pubB64 := base64.StdEncoding.EncodeToString(pub)
deviceId := "device-test-1"
fp := strings.Repeat("a", 64)
return &model.Token{
Id: 1,
UserId: 42,
DeviceId: &deviceId,
DevicePubkey: &pubB64,
DeviceFingerprint: fp,
}
}
func TestVerifyDeviceSignatureIfRequired(t *testing.T) {
gin.SetMode(gin.TestMode)
priv := loadRFC8032PrivateKey(t)
tests := []struct {
name string
setup func() (*gin.Context, *model.Token)
wantErrIs error
}{
{
name: "happy_path_signed_request",
setup: func() (*gin.Context, *model.Token) {
r := defaultSignedReq()
return makeSignedContext(t, r, priv), makeBoundToken(priv)
},
wantErrIs: nil,
},
{
name: "expired_timestamp",
setup: func() (*gin.Context, *model.Token) {
r := defaultSignedReq()
r.Timestamp = time.Now().UnixMilli() - 10*60*1000 // 10 min ago
return makeSignedContext(t, r, priv), makeBoundToken(priv)
},
wantErrIs: ErrDeviceTimestampOutOfWindow,
},
{
name: "future_timestamp_beyond_window",
setup: func() (*gin.Context, *model.Token) {
r := defaultSignedReq()
r.Timestamp = time.Now().UnixMilli() + 10*60*1000 // 10 min ahead
return makeSignedContext(t, r, priv), makeBoundToken(priv)
},
wantErrIs: ErrDeviceTimestampOutOfWindow,
},
{
name: "fingerprint_mismatch",
setup: func() (*gin.Context, *model.Token) {
r := defaultSignedReq()
c := makeSignedContext(t, r, priv)
// Server stored a different fingerprint than the client reports
tok := makeBoundToken(priv)
tok.DeviceFingerprint = strings.Repeat("b", 64)
return c, tok
},
wantErrIs: ErrDeviceFingerprintMismatch,
},
{
name: "wrong_signature",
setup: func() (*gin.Context, *model.Token) {
r := defaultSignedReq()
c := makeSignedContext(t, r, priv)
// Tamper one byte of signature
orig := c.Request.Header.Get(HeaderSignature)
raw, _ := base64.StdEncoding.DecodeString(orig)
raw[0] ^= 0xff
c.Request.Header.Set(HeaderSignature, base64.StdEncoding.EncodeToString(raw))
return c, makeBoundToken(priv)
},
wantErrIs: ErrDeviceSignatureInvalid,
},
{
name: "tampered_body",
setup: func() (*gin.Context, *model.Token) {
r := defaultSignedReq()
c := makeSignedContext(t, r, priv)
// Replace request body AFTER signing → bodyHash will no longer match
c.Request.Body = http.NoBody
c.Request.ContentLength = 0
return c, makeBoundToken(priv)
},
wantErrIs: ErrDeviceSignatureInvalid,
},
{
name: "missing_signature_header_on_bound_token",
setup: func() (*gin.Context, *model.Token) {
r := defaultSignedReq()
c := makeSignedContext(t, r, priv)
c.Request.Header.Del(HeaderSignature)
return c, makeBoundToken(priv)
},
wantErrIs: ErrDeviceSignatureRequired,
},
{
name: "device_id_mismatch",
setup: func() (*gin.Context, *model.Token) {
r := defaultSignedReq()
c := makeSignedContext(t, r, priv)
c.Request.Header.Set(HeaderDeviceID, "different-device")
return c, makeBoundToken(priv)
},
wantErrIs: ErrDeviceSignatureInvalid,
},
{
name: "revoked_token",
setup: func() (*gin.Context, *model.Token) {
r := defaultSignedReq()
c := makeSignedContext(t, r, priv)
tok := makeBoundToken(priv)
tok.RevokedAt = time.Now().UnixMilli()
return c, tok
},
wantErrIs: ErrDeviceRevoked,
},
{
name: "legacy_token_no_binding_no_enforcement",
setup: func() (*gin.Context, *model.Token) {
r := defaultSignedReq()
c := makeSignedContext(t, r, priv)
// Strip signature headers — looks like a legacy sk- caller
c.Request.Header.Del(HeaderSignature)
c.Request.Header.Del(HeaderDeviceID)
c.Request.Header.Del(HeaderTimestamp)
c.Request.Header.Del(HeaderNonce)
c.Request.Header.Del(HeaderFingerprint)
tok := &model.Token{Id: 99, UserId: 42}
return c, tok
},
wantErrIs: nil,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
// Ensure global enforcement off for the legacy_token case.
operation_setting.GetDeviceBindingSetting().RequireGlobal = false
c, tok := tc.setup()
err := VerifyDeviceSignatureIfRequired(c, tok)
if tc.wantErrIs == nil {
if err != nil {
t.Fatalf("expected ok, got error: %v", err)
}
return
}
if !errors.Is(err, tc.wantErrIs) {
t.Fatalf("expected %v, got %v", tc.wantErrIs, err)
}
})
}
}
func TestNonceReplayDetection(t *testing.T) {
gin.SetMode(gin.TestMode)
priv := loadRFC8032PrivateKey(t)
// Use a freshly built request to keep timestamp fresh
r := defaultSignedReq()
r.Nonce = randomNonce()
tok := makeBoundToken(priv)
// First request must succeed
c1 := makeSignedContext(t, r, priv)
if err := VerifyDeviceSignatureIfRequired(c1, tok); err != nil {
t.Fatalf("first request: expected ok, got %v", err)
}
// Same nonce + device_id replayed → must reject as replay
c2 := makeSignedContext(t, r, priv)
err := VerifyDeviceSignatureIfRequired(c2, tok)
if !errors.Is(err, ErrDeviceNonceReplayed) {
t.Fatalf("replay: expected ErrDeviceNonceReplayed, got %v", err)
}
}
// TestCanonicalStringMatchesVectors loads the cross-language test vectors
// and asserts the Go implementation produces the documented canonical
// string byte-for-byte. The Rust + TS tests load the same file and run
// the equivalent assertion; if either drifts, both fail loudly.
func TestCanonicalStringMatchesVectors(t *testing.T) {
path := filepath.Join("..", "testdata", "device_signature_vectors.json")
data, err := os.ReadFile(path)
if err != nil {
t.Skipf("vectors file not readable from this working dir: %v", err)
}
var doc struct {
Cases []struct {
Name string `json:"name"`
Input struct {
Method string `json:"method"`
Path string `json:"path_with_query"`
TimestampMs string `json:"timestamp_ms"`
Nonce string `json:"nonce_hex"`
Fingerprint string `json:"device_fingerprint"`
BodyText string `json:"body_text"`
} `json:"input"`
ExpectedBodyHash string `json:"expected_body_sha256_hex"`
ExpectedCanonical string `json:"expected_canonical"`
ExpectedCanonicalStarts string `json:"expected_canonical_starts_with"`
} `json:"cases"`
}
if err := json.Unmarshal(data, &doc); err != nil {
t.Fatalf("unmarshal vectors: %v", err)
}
for _, tc := range doc.Cases {
t.Run(tc.Name, func(t *testing.T) {
bodySum := sha256.Sum256([]byte(tc.Input.BodyText))
gotBodyHash := hex.EncodeToString(bodySum[:])
if tc.ExpectedBodyHash != "" && tc.ExpectedBodyHash != gotBodyHash {
// Some vectors use illustrative body_hash that may not
// match — only fail when the vector file claims a real
// expected value (the empty-body case).
if tc.Name == "GET_empty_body" {
t.Fatalf("body hash mismatch for %s: got %s want %s",
tc.Name, gotBodyHash, tc.ExpectedBodyHash)
}
}
canonical := strings.Join([]string{
strings.ToUpper(tc.Input.Method),
tc.Input.Path,
tc.Input.TimestampMs,
tc.Input.Nonce,
tc.Input.Fingerprint,
gotBodyHash,
}, "\n")
if tc.ExpectedCanonical != "" && tc.ExpectedCanonical != canonical {
t.Fatalf("canonical mismatch for %s\n got: %q\n want: %q",
tc.Name, canonical, tc.ExpectedCanonical)
}
if tc.ExpectedCanonicalStarts != "" && !strings.HasPrefix(canonical, tc.ExpectedCanonicalStarts) {
t.Fatalf("canonical prefix mismatch for %s\n got: %q\n want prefix: %q",
tc.Name, canonical, tc.ExpectedCanonicalStarts)
}
})
}
}