tailer: require pinned log keys and stop on STH verification failure

VerifySTH silently passed when no log key was pinned or when an STH
arrived without a signature, and a failed verification only produced a
log line while ingestion continued from the unauthenticated source.
NewClient now refuses logs without keys, the tailer skips the fetch
cycle on verification failure, and it rejects list entries whose
log_id does not match SHA-256 of the key.
This commit is contained in:
lakshit verma 2026-08-21 21:30:40 +05:30
parent 394174763d
commit 250652a3dd
No known key found for this signature in database
3 changed files with 84 additions and 24 deletions

View file

@ -36,19 +36,20 @@ type Client struct {
func NewClient(baseURL, keyB64 string) (*Client, error) { func NewClient(baseURL, keyB64 string) (*Client, error) {
c := &Client{BaseURL: baseURL, HTTP: &http.Client{Timeout: 30 * time.Second}} c := &Client{BaseURL: baseURL, HTTP: &http.Client{Timeout: 30 * time.Second}}
if keyB64 != "" { if keyB64 == "" {
der, err := base64.StdEncoding.DecodeString(keyB64) return nil, fmt.Errorf("missing log key")
if err != nil {
return nil, fmt.Errorf("bad log key: %w", err)
}
pub, err := x509.ParsePKIXPublicKey(der)
if err != nil {
return nil, fmt.Errorf("bad log key: %w", err)
}
c.pubKey = pub
sum := sha256.Sum256(der)
c.logID = sum[:]
} }
der, err := base64.StdEncoding.DecodeString(keyB64)
if err != nil {
return nil, fmt.Errorf("bad log key: %w", err)
}
pub, err := x509.ParsePKIXPublicKey(der)
if err != nil {
return nil, fmt.Errorf("bad log key: %w", err)
}
c.pubKey = pub
sum := sha256.Sum256(der)
c.logID = sum[:]
return c, nil return c, nil
} }
@ -114,8 +115,11 @@ func (e *HTTPError) Error() string {
} }
func (c *Client) VerifySTH(sth *STH) error { func (c *Client) VerifySTH(sth *STH) error {
if c.pubKey == nil || len(sth.TreeHeadSignature) == 0 { if c.pubKey == nil {
return nil return fmt.Errorf("no log key pinned")
}
if len(sth.TreeHeadSignature) == 0 {
return fmt.Errorf("sth has no signature")
} }
ds, err := decodeDigitallySigned(sth.TreeHeadSignature) ds, err := decodeDigitallySigned(sth.TreeHeadSignature)
if err != nil { if err != nil {

View file

@ -62,6 +62,10 @@ func (t *Tailer) tailOne(ctx context.Context, lg loglist.Log, state string) {
log.Printf("tail %s: %v", lg.LogID, err) log.Printf("tail %s: %v", lg.LogID, err)
return return
} }
if expected := client.LogID(); lg.LogID != "" && expected != lg.LogID {
log.Printf("tail %s: log list key mismatch (sha256(key) = %s), skipping", lg.LogID, expected)
return
}
id := lg.LogID id := lg.LogID
wm, err := t.Store.Watermark(id) wm, err := t.Store.Watermark(id)
if err != nil { if err != nil {
@ -105,6 +109,10 @@ func (t *Tailer) tailOne(ctx context.Context, lg loglist.Log, state string) {
} }
if err := client.VerifySTH(sth); err != nil { if err := client.VerifySTH(sth); err != nil {
log.Printf("tail %s: verify: %v", client.ShortID(), err) log.Printf("tail %s: verify: %v", client.ShortID(), err)
if !sleepCtx(ctx, jitter(interval)) {
return
}
continue
} }
target := sth.TreeSize target := sth.TreeSize
if drainTarget >= 0 && drainTarget < target { if drainTarget >= 0 && drainTarget < target {

View file

@ -2,8 +2,11 @@ package tailer
import ( import (
"context" "context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand" "crypto/rand"
"crypto/rsa" "crypto/rsa"
"crypto/sha256"
"crypto/x509" "crypto/x509"
"encoding/base64" "encoding/base64"
"encoding/binary" "encoding/binary"
@ -24,6 +27,45 @@ type fakeLog struct {
entries [][]byte entries [][]byte
sthHits int sthHits int
shortRead bool shortRead bool
key *ecdsa.PrivateKey
}
func newFakeLog(t *testing.T) *fakeLog {
t.Helper()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatal(err)
}
return &fakeLog{key: key}
}
// keyB64 returns the log's public key in the log list format, and logID its
// RFC 6962 log id (base64 of SHA-256 of the public key DER).
func (f *fakeLog) keyB64(t *testing.T) (keyB64, logID string) {
t.Helper()
der, err := x509.MarshalPKIXPublicKey(&f.key.PublicKey)
if err != nil {
t.Fatal(err)
}
sum := sha256.Sum256(der)
return base64.StdEncoding.EncodeToString(der), base64.StdEncoding.EncodeToString(sum[:])
}
func (f *fakeLog) signSTH(t *testing.T, ts int64, size int64, root []byte) []byte {
t.Helper()
input := make([]byte, 0, 2+8+8+32)
input = append(input, 0, 1)
input = binary.BigEndian.AppendUint64(input, uint64(ts))
input = binary.BigEndian.AppendUint64(input, uint64(size))
input = append(input, root...)
digest := sha256.Sum256(input)
sig, err := ecdsa.SignASN1(rand.Reader, f.key, digest[:])
if err != nil {
t.Fatal(err)
}
out := []byte{4, 3}
out = binary.BigEndian.AppendUint16(out, uint16(len(sig)))
return append(out, sig...)
} }
func (f *fakeLog) leafInput(der []byte, ts int64) []byte { func (f *fakeLog) leafInput(der []byte, ts int64) []byte {
@ -45,9 +87,10 @@ func (f *fakeLog) serve(t *testing.T) *httptest.Server {
root := make([]byte, 32) root := make([]byte, 32)
rand.Read(root) rand.Read(root)
json.NewEncoder(w).Encode(map[string]any{ json.NewEncoder(w).Encode(map[string]any{
"tree_size": len(f.entries), "tree_size": len(f.entries),
"timestamp": time.Now().UnixMilli(), "timestamp": time.Now().UnixMilli(),
"sha256_root_hash": base64.StdEncoding.EncodeToString(root), "sha256_root_hash": base64.StdEncoding.EncodeToString(root),
"tree_head_signature": base64.StdEncoding.EncodeToString(f.signSTH(t, time.Now().UnixMilli(), int64(len(f.entries)), root)),
}) })
case "/ct/v1/get-entries": case "/ct/v1/get-entries":
var start, end int64 var start, end int64
@ -99,7 +142,9 @@ func TestTailToSearchRoundTrip(t *testing.T) {
for _, short := range []bool{false, true} { for _, short := range []bool{false, true} {
name := fmt.Sprintf("shortread=%v", short) name := fmt.Sprintf("shortread=%v", short)
t.Run(name, func(t *testing.T) { t.Run(name, func(t *testing.T) {
fl := &fakeLog{shortRead: short} fl := newFakeLog(t)
fl.shortRead = short
keyB64, logID := fl.keyB64(t)
ts1 := time.Date(2020, 1, 1, 0, 0, 0, 0, time.UTC).UnixMilli() ts1 := time.Date(2020, 1, 1, 0, 0, 0, 0, time.UTC).UnixMilli()
ts2 := time.Date(2021, 6, 15, 12, 0, 0, 0, time.UTC).UnixMilli() ts2 := time.Date(2021, 6, 15, 12, 0, 0, 0, time.UTC).UnixMilli()
c1 := makeLeafCert(t, []string{"roundtrip.example.com", "www.roundtrip.example.com"}) c1 := makeLeafCert(t, []string{"roundtrip.example.com", "www.roundtrip.example.com"})
@ -116,7 +161,8 @@ func TestTailToSearchRoundTrip(t *testing.T) {
tr := &Tailer{Store: st, Interval: time.Millisecond, Window: 1} tr := &Tailer{Store: st, Interval: time.Millisecond, Window: 1}
logs := []loglist.Log{{ logs := []loglist.Log{{
LogID: "testlog", LogID: logID,
Key: keyB64,
URL: srv.URL, URL: srv.URL,
State: map[string]loglist.StateDetail{"usable": {}}, State: map[string]loglist.StateDetail{"usable": {}},
}} }}
@ -149,7 +195,7 @@ func TestTailToSearchRoundTrip(t *testing.T) {
if res[0].Sub != "roundtrip.example.com" || res[2].Sub != "mail.roundtrip.example.com" { if res[0].Sub != "roundtrip.example.com" || res[2].Sub != "mail.roundtrip.example.com" {
t.Errorf("insertion order wrong: %+v", res) t.Errorf("insertion order wrong: %+v", res)
} }
wm, err := st.Watermark("testlog") wm, err := st.Watermark(logID)
if err != nil || wm != 2 { if err != nil || wm != 2 {
t.Errorf("watermark = %d, %v", wm, err) t.Errorf("watermark = %d, %v", wm, err)
} }
@ -158,7 +204,8 @@ func TestTailToSearchRoundTrip(t *testing.T) {
} }
func TestSTHRegressionIgnored(t *testing.T) { func TestSTHRegressionIgnored(t *testing.T) {
fl := &fakeLog{} fl := newFakeLog(t)
keyB64, logID := fl.keyB64(t)
der := makeLeafCert(t, []string{"reg.example.com"}) der := makeLeafCert(t, []string{"reg.example.com"})
fl.entries = [][]byte{fl.leafInput(der, 1000)} fl.entries = [][]byte{fl.leafInput(der, 1000)}
srv := fl.serve(t) srv := fl.serve(t)
@ -169,17 +216,18 @@ func TestSTHRegressionIgnored(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
defer st.Close() defer st.Close()
if err := st.SetWatermarkSync("reglog", 5); err != nil { if err := st.SetWatermarkSync(logID, 5); err != nil {
t.Fatal(err) t.Fatal(err)
} }
tr := &Tailer{Store: st, Interval: time.Millisecond} tr := &Tailer{Store: st, Interval: time.Millisecond}
tr.Sync(context.Background(), []loglist.Log{{ tr.Sync(context.Background(), []loglist.Log{{
LogID: "reglog", LogID: logID,
Key: keyB64,
URL: srv.URL, URL: srv.URL,
State: map[string]loglist.StateDetail{"usable": {}}, State: map[string]loglist.StateDetail{"usable": {}},
}}) }})
time.Sleep(50 * time.Millisecond) time.Sleep(50 * time.Millisecond)
wm, _ := st.Watermark("reglog") wm, _ := st.Watermark(logID)
if wm < 5 { if wm < 5 {
t.Errorf("watermark rewound to %d", wm) t.Errorf("watermark rewound to %d", wm)
} }