diff --git a/client.go b/client.go index 9e32217..5bb1dc4 100644 --- a/client.go +++ b/client.go @@ -431,13 +431,17 @@ func (c *Client) RevokeLocalVaultInstallation(ctx context.Context, installationI func (c *Client) RelayLocalVaultEvent(ctx context.Context, installationID string, envelope M, payload M, envelopeHash, signatureHex, publicKeyHex string) (*EventReceipt, error) { body := M{"envelope": envelope, "payload": payload} + raw, err := json.Marshal(body) + if err != nil { + return nil, fmt.Errorf("marshal local vault event: %w", err) + } headers := map[string]string{ "X-Attesto-Local-Vault-Envelope-Hash": envelopeHash, "X-Attesto-Local-Vault-Signature": signatureHex, "X-Attesto-Local-Vault-Public-Key": publicKeyHex, } var out EventReceipt - err := c.requestRaw(ctx, http.MethodPost, "/v2/local-vault/installations/"+url.PathEscape(installationID)+"/events", nil, mustJSON(body), headers, "", &out) + err = c.requestRaw(ctx, http.MethodPost, "/v2/local-vault/installations/"+url.PathEscape(installationID)+"/events", nil, raw, headers, "", &out) return &out, err } @@ -650,14 +654,6 @@ func sleep(attempt int) { time.Sleep(time.Duration(100*attempt) * time.Millisecond) } -func mustJSON(value any) []byte { - raw, err := json.Marshal(value) - if err != nil { - panic(err) - } - return raw -} - func sanitizeMessage(message string) string { for _, marker := range []string{"sk_live_", "sk_test_", "pk_live_", "pk_test_", "npm_", "pypi-", "atto_live_", "atto_test_"} { if strings.Contains(message, marker) { diff --git a/client_test.go b/client_test.go index a8c2a1b..2427e20 100644 --- a/client_test.go +++ b/client_test.go @@ -105,3 +105,33 @@ func TestBearerClientCanCallTenantEndpoints(t *testing.T) { t.Fatalf("unexpected streams: %#v", streams) } } + +func TestRelayLocalVaultEventReturnsMarshalError(t *testing.T) { + called := false + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + called = true + t.Fatalf("request should not be sent when local vault event body cannot be marshaled") + })) + defer server.Close() + + client, err := NewClient(testAPIKey, WithBaseURL(server.URL), WithMaxRetries(1)) + if err != nil { + t.Fatalf("client: %v", err) + } + + _, err = client.RelayLocalVaultEvent( + context.Background(), + "lv_123", + M{"source": "local-vault"}, + M{"bad": func() {}}, + strings.Repeat("a", 64), + strings.Repeat("b", 128), + strings.Repeat("c", 64), + ) + if err == nil || !strings.Contains(err.Error(), "marshal local vault event") { + t.Fatalf("expected marshal error, got %v", err) + } + if called { + t.Fatalf("request was sent after marshal failure") + } +}