diff --git a/client.go b/client.go index 5bb1dc4..fb78bfa 100644 --- a/client.go +++ b/client.go @@ -29,6 +29,7 @@ type Client struct { type Option func(*Client) error var apiKeyPattern = regexp.MustCompile(`^atto_(?:live|test)_[0-9a-f]{32}$`) +var secureRandomRead = rand.Read func NewClient(apiKey string, opts ...Option) (*Client, error) { if !apiKeyPattern.MatchString(apiKey) { @@ -129,7 +130,7 @@ func (c *Client) CreateStream(ctx context.Context, input StreamCreateInput, opti if input.Metadata == nil { input.Metadata = M{} } - err := c.requestJSON(ctx, http.MethodPost, "/v2/streams", nil, input, idempotency(options), &out) + err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/streams", nil, input, &out, options) return &out, err } @@ -184,7 +185,7 @@ func (c *Client) LogEvent(ctx context.Context, streamID string, input EventInput } } var out EventReceipt - if err := c.requestJSON(ctx, http.MethodPost, "/v2/streams/"+url.PathEscape(streamID)+"/events", nil, input, idempotency(options), &out); err != nil { + if err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/streams/"+url.PathEscape(streamID)+"/events", nil, input, &out, options); err != nil { return nil, err } if err := c.trackHead(out); err != nil { @@ -223,7 +224,7 @@ func (c *Client) LogEvents(ctx context.Context, streamID string, events []EventI } body := M{"events": events} var out EventBatchResponse - if err := c.requestJSON(ctx, http.MethodPost, "/v2/streams/"+url.PathEscape(streamID)+"/events/batch", nil, body, idempotency(options), &out); err != nil { + if err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/streams/"+url.PathEscape(streamID)+"/events/batch", nil, body, &out, options); err != nil { return nil, err } for _, receipt := range out.Receipts { @@ -270,19 +271,19 @@ func (c *Client) GetIVCEpoch(ctx context.Context, ivcEpochID string) (M, error) func (c *Client) BuildVerifierBundle(ctx context.Context, fromCheckpointID, toCheckpointID string, options ...RequestOptions) (*VerifierBundle, error) { body := M{"fromCheckpointId": fromCheckpointID, "toCheckpointId": toCheckpointID} var out VerifierBundle - err := c.requestJSON(ctx, http.MethodPost, "/v2/audit/packs", nil, body, idempotency(options), &out) + err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/audit/packs", nil, body, &out, options) return &out, err } func (c *Client) VerifyReceiptRemote(ctx context.Context, input ReceiptVerifyInput, options ...RequestOptions) (*VerifyReport, error) { var out VerifyReport - err := c.requestJSON(ctx, http.MethodPost, "/v2/verify/receipt", nil, input, idempotency(options), &out) + err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/verify/receipt", nil, input, &out, options) return &out, err } func (c *Client) VerifyObjectRemote(ctx context.Context, input OfflineVerifyInput, options ...RequestOptions) (*VerifyReport, error) { var out VerifyReport - err := c.requestJSON(ctx, http.MethodPost, "/v2/verify", nil, input, idempotency(options), &out) + err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/verify", nil, input, &out, options) return &out, err } @@ -313,7 +314,7 @@ func (c *Client) ListTenantIVCEpochs(ctx context.Context, streamID string, limit func (c *Client) BuildTenantAuditPack(ctx context.Context, fromCheckpointID, toCheckpointID string, options ...RequestOptions) (*VerifierBundle, error) { body := M{"fromCheckpointId": fromCheckpointID, "toCheckpointId": toCheckpointID} var out VerifierBundle - err := c.requestJSON(ctx, http.MethodPost, "/v2/tenant/audit/packs", nil, body, idempotency(options), &out) + err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/tenant/audit/packs", nil, body, &out, options) return &out, err } @@ -323,7 +324,7 @@ func (c *Client) ListSignedWebhookConnectors(ctx context.Context, limit, offset func (c *Client) CreateSignedWebhookConnector(ctx context.Context, input ConnectorCreateInput, options ...RequestOptions) (*Connector, error) { var out Connector - err := c.requestJSON(ctx, http.MethodPost, "/v2/tenant/connectors/signed-webhooks", nil, input, idempotency(options), &out) + err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/tenant/connectors/signed-webhooks", nil, input, &out, options) return &out, err } @@ -337,13 +338,13 @@ func (c *Client) ListS3ObjectConnectors(ctx context.Context, limit, offset int) func (c *Client) CreateS3ObjectConnector(ctx context.Context, input S3ConnectorCreateInput, options ...RequestOptions) (*Connector, error) { var out Connector - err := c.requestJSON(ctx, http.MethodPost, "/v2/tenant/connectors/s3-objects", nil, input, idempotency(options), &out) + err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/tenant/connectors/s3-objects", nil, input, &out, options) return &out, err } func (c *Client) CommitS3Object(ctx context.Context, connectorID string, body M, options ...RequestOptions) (*EventReceipt, error) { var out EventReceipt - err := c.requestJSON(ctx, http.MethodPost, "/v2/tenant/connectors/s3-objects/"+url.PathEscape(connectorID)+"/commit", nil, body, idempotency(options), &out) + err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/tenant/connectors/s3-objects/"+url.PathEscape(connectorID)+"/commit", nil, body, &out, options) return &out, err } @@ -357,7 +358,7 @@ func (c *Client) ListRepositoryWebhookConnectors(ctx context.Context, limit, off func (c *Client) CreateRepositoryWebhookConnector(ctx context.Context, input RepositoryConnectorCreateInput, options ...RequestOptions) (*Connector, error) { var out Connector - err := c.requestJSON(ctx, http.MethodPost, "/v2/tenant/connectors/repository-webhooks", nil, input, idempotency(options), &out) + err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/tenant/connectors/repository-webhooks", nil, input, &out, options) return &out, err } @@ -421,7 +422,7 @@ func (c *Client) ListLocalVaultInstallations(ctx context.Context, limit, offset func (c *Client) CreateLocalVaultInstallation(ctx context.Context, input LocalVaultInstallationCreateInput, options ...RequestOptions) (*LocalVaultInstallation, error) { var out LocalVaultInstallation - err := c.requestJSON(ctx, http.MethodPost, "/v2/tenant/local-vault/installations", nil, input, idempotency(options), &out) + err := c.requestJSONIdempotent(ctx, http.MethodPost, "/v2/tenant/local-vault/installations", nil, input, &out, options) return &out, err } @@ -446,11 +447,11 @@ func (c *Client) RelayLocalVaultEvent(ctx context.Context, installationID string } func (c *Client) SubmitLocalVaultWitnessReceipt(ctx context.Context, installationID string, receipt M, options ...RequestOptions) (M, error) { - return c.postObject(ctx, "/v2/local-vault/installations/"+url.PathEscape(installationID)+"/witness/checkpoints", M{"receipt": receipt}, idempotency(options)) + return c.postObjectIdempotent(ctx, "/v2/local-vault/installations/"+url.PathEscape(installationID)+"/witness/checkpoints", M{"receipt": receipt}, options) } func (c *Client) SubmitLocalVaultForkEvidence(ctx context.Context, installationID string, forkEvidence M, options ...RequestOptions) (M, error) { - return c.postObject(ctx, "/v2/local-vault/installations/"+url.PathEscape(installationID)+"/witness/checkpoints", M{"forkEvidence": forkEvidence}, idempotency(options)) + return c.postObjectIdempotent(ctx, "/v2/local-vault/installations/"+url.PathEscape(installationID)+"/witness/checkpoints", M{"forkEvidence": forkEvidence}, options) } func (c *Client) GetMarketplaceItem(ctx context.Context, slug string) (M, error) { @@ -458,13 +459,13 @@ func (c *Client) GetMarketplaceItem(ctx context.Context, slug string) (M, error) } func (c *Client) SubmitMarketplaceAsset(ctx context.Context, input MarketplaceAssetSubmitInput, options ...RequestOptions) (M, error) { - return c.postObject(ctx, "/v1/marketplace/publisher/assets", M{ + return c.postObjectIdempotent(ctx, "/v1/marketplace/publisher/assets", M{ "manifest": input.Manifest, "sourceRef": input.SourceRef, "visibility": input.Visibility, "pricingModel": input.PricingModel, "priceCents": input.PriceCents, - }, idempotency(options)) + }, options) } func (c *Client) ListMarketplaceReviewAssets(ctx context.Context, state string) ([]M, error) { @@ -478,7 +479,7 @@ func (c *Client) ListMarketplaceReviewAssets(ctx context.Context, state string) } func (c *Client) ApproveMarketplaceAsset(ctx context.Context, slug string, reason string, options ...RequestOptions) (M, error) { - return c.postObject(ctx, "/v1/platform/marketplace/assets/"+url.PathEscape(slug)+"/approve", M{"reason": reason}, idempotency(options)) + return c.postObjectIdempotent(ctx, "/v1/platform/marketplace/assets/"+url.PathEscape(slug)+"/approve", M{"reason": reason}, options) } func (c *Client) getObject(ctx context.Context, path string, values url.Values) (M, error) { @@ -493,6 +494,14 @@ func (c *Client) postObject(ctx context.Context, path string, body M, idempotenc return out, err } +func (c *Client) postObjectIdempotent(ctx context.Context, path string, body M, options []RequestOptions) (M, error) { + idempotencyKey, err := idempotency(options) + if err != nil { + return nil, err + } + return c.postObject(ctx, path, body, idempotencyKey) +} + func (c *Client) getList(ctx context.Context, path string, limit, offset int) ([]M, error) { values := url.Values{} setPaging(values, limit, offset) @@ -525,6 +534,14 @@ func (c *Client) requestJSON(ctx context.Context, method, path string, values ur return c.requestRaw(ctx, method, path, values, raw, nil, idempotencyKey, out) } +func (c *Client) requestJSONIdempotent(ctx context.Context, method, path string, values url.Values, body any, out any, options []RequestOptions) error { + idempotencyKey, err := idempotency(options) + if err != nil { + return err + } + return c.requestJSON(ctx, method, path, values, body, idempotencyKey, out) +} + func (c *Client) requestRaw(ctx context.Context, method, path string, values url.Values, body []byte, extraHeaders map[string]string, idempotencyKey string, out any) error { if values != nil && len(values) > 0 { path += "?" + values.Encode() @@ -630,15 +647,15 @@ func skipPreflight(options []RequestOptions) bool { return len(options) > 0 && options[0].SkipPreflight } -func idempotency(options []RequestOptions) string { +func idempotency(options []RequestOptions) (string, error) { if len(options) > 0 && options[0].IdempotencyKey != "" { - return options[0].IdempotencyKey + return options[0].IdempotencyKey, nil } var raw [16]byte - if _, err := rand.Read(raw[:]); err != nil { - return fmt.Sprintf("%d", time.Now().UnixNano()) + if _, err := secureRandomRead(raw[:]); err != nil { + return "", fmt.Errorf("generate idempotency key: %w", err) } - return hex.EncodeToString(raw[:]) + return hex.EncodeToString(raw[:]), nil } func setPaging(values url.Values, limit, offset int) { diff --git a/client_test.go b/client_test.go index 2427e20..ce16066 100644 --- a/client_test.go +++ b/client_test.go @@ -3,6 +3,7 @@ package attesto import ( "context" "encoding/json" + "errors" "net/http" "net/http/httptest" "strings" @@ -135,3 +136,67 @@ func TestRelayLocalVaultEventReturnsMarshalError(t *testing.T) { t.Fatalf("request was sent after marshal failure") } } + +func TestGeneratedIdempotencyKeyFailsClosedOnEntropyError(t *testing.T) { + originalRead := secureRandomRead + secureRandomRead = func([]byte) (int, error) { + return 0, errors.New("entropy unavailable") + } + t.Cleanup(func() { secureRandomRead = originalRead }) + + called := false + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + called = true + t.Fatalf("request should not be sent when idempotency key generation fails") + })) + defer server.Close() + + client, err := NewClient(testAPIKey, WithBaseURL(server.URL), WithMaxRetries(1)) + if err != nil { + t.Fatalf("client: %v", err) + } + + _, err = client.CreateStream(context.Background(), StreamCreateInput{UseCase: "ai-governance", PolicyID: "policy-main"}) + if err == nil || !strings.Contains(err.Error(), "generate idempotency key") { + t.Fatalf("expected idempotency entropy error, got %v", err) + } + if called { + t.Fatalf("request was sent after idempotency generation failure") + } +} + +func TestExplicitIdempotencyKeyBypassesEntropyGeneration(t *testing.T) { + originalRead := secureRandomRead + secureRandomRead = func([]byte) (int, error) { + return 0, errors.New("entropy unavailable") + } + t.Cleanup(func() { secureRandomRead = originalRead }) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Idempotency-Key") != "fixed-key" { + t.Fatalf("explicit idempotency key missing: %q", r.Header.Get("Idempotency-Key")) + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(Stream{ + StreamID: "str_fixed", SystemID: "sys_fixed", UseCase: "ai-governance", PolicyID: "policy-main", Status: "active", Created: true, + }) + })) + defer server.Close() + + client, err := NewClient(testAPIKey, WithBaseURL(server.URL), WithMaxRetries(1)) + if err != nil { + t.Fatalf("client: %v", err) + } + + stream, err := client.CreateStream( + context.Background(), + StreamCreateInput{UseCase: "ai-governance", PolicyID: "policy-main"}, + RequestOptions{IdempotencyKey: "fixed-key"}, + ) + if err != nil { + t.Fatalf("create stream with explicit idempotency key: %v", err) + } + if stream.StreamID != "str_fixed" { + t.Fatalf("unexpected stream: %#v", stream) + } +}