Files
chenchen 80d0f956d0 feat(manager): P0 device-binding signature layer for cc-haha desktop clients
Lays down the server side of a per-request Ed25519 signature scheme that
binds a token to a specific desktop install, so the bearer key can't be
extracted from ~/.claude/cc-haha/providers.json and resold. Plan lives
at ~/.claude/plans/peaceful-sprouting-crane.md.

Compatibility: legacy bare-bearer sk- callers (CLI/SDK) pass through
unchanged until P3 (30-day deadline) flips RequireGlobal=true. No
existing token rows are modified — pubkey is nullable and defaults to
null.

Pieces:
- model.Token gains DeviceId, DevicePubkey, DeviceFingerprint, DeviceName,
  DevicePlatform, DeviceAppVersion, DeviceBoundAt, DeviceLastSeenIp,
  DeviceLastUsedAt, RequireDeviceBinding, RevokedAt, RevokedReason.
  Pure additive columns, GORM AutoMigrate handles SQLite/MySQL/PG.
- common.VerifyEd25519Signature: thin wrapper around crypto/ed25519
  stdlib, used by the new middleware. No new external deps.
- service.MarkNonceUsed: Redis SETNX-based nonce store with an
  in-memory sync.Map fallback for single-instance dev. TTL = setting.
- middleware.VerifyDeviceSignatureIfRequired: wired into TokenAuth as
  a fail-fast dispatch right after model.ValidateUserToken. Verifies
  canonical = METHOD\nPATH\nTS_MS\nNONCE\nFINGERPRINT\nSHA256(BODY),
  signed as Ed25519(sha256(canonical)). 120s timestamp window, 300s
  nonce window, fingerprint stored at pair time must match the header.
- controller.PairDevice / ListUserDevices / RenameUserDevice /
  RevokeUserDevice, mounted at /api/devices/* behind UserAuth().
  PairDevice enforces 5-per-user cap and returns the raw sk- once,
  to be stored in the client's OS keychain (not providers.json).
- operation_setting.DeviceBindingSetting: MaxDevicesPerUser=5,
  TimestampWindowMs=120000, NonceTTLSec=300, RequireGlobal=false.

Tests:
- common/crypto_test.go covers round-trip + tamper + malformed inputs.
- middleware/device_signature_test.go covers all error-path branches
  (expired ts, wrong sig, tampered body, fingerprint mismatch, replay,
  revoked, missing headers, legacy fallthrough).
- testdata/device_signature_vectors.json is the cross-language contract
  Rust+TS sides will load to assert byte-identical canonical strings.

Untouched but reserved for follow-up phases:
- Anomaly detection / IP-diversity flagging (P1)
- 30-day deprecation banner + email notifications (P2)
- Hard cutover RequireGlobal=true (P3, day 31)
2026-05-20 12:10:57 +08:00

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)
}
})
}
}