diff --git a/server/cmd/gram/streams.go b/server/cmd/gram/streams.go index f9ef81494d4..87f2e5a736e 100644 --- a/server/cmd/gram/streams.go +++ b/server/cmd/gram/streams.go @@ -516,7 +516,6 @@ func newStreamsCommand() *cli.Command { replicaDB, encryptionClient, guardianPolicy, - featureFlags, ) metricRelayHandler := otelsvc.NewMetricRelayHandler( diff --git a/server/internal/dataexports/queries.sql b/server/internal/dataexports/queries.sql index 87bea66537d..1110d0bb285 100644 --- a/server/internal/dataexports/queries.sql +++ b/server/internal/dataexports/queries.sql @@ -1,5 +1,25 @@ -- Every query pins both organization_id and project_id. Resource UUIDs and -- tenant-pinned foreign keys are integrity controls, not authorization bounds. + +-- Resolve the active OTEL destination for one project data source. +-- name: GetActiveOtelRouteDestination :one +SELECT + destination.endpoint_url, + destination.headers_encrypted, + COALESCE(destination.sensitive_data, 'exclude') = 'include' AS include_sensitive_data +FROM data_export_routes AS route +JOIN otel_destinations AS destination + ON destination.organization_id = route.organization_id + AND destination.project_id = route.project_id + AND destination.id = route.otel_destination_id +WHERE route.organization_id = @organization_id + AND route.project_id = @project_id + AND route.data_source = @data_source + AND route.enabled IS TRUE + AND route.deleted IS FALSE + AND route.otel_destination_id IS NOT NULL + AND destination.deleted IS FALSE; + -- List active destinations in stable creation order for the management API. -- name: ListOtelDestinations :many SELECT * diff --git a/server/internal/dataexports/repo/queries.sql.go b/server/internal/dataexports/repo/queries.sql.go index 5e9cfad6224..75bd1efc61e 100644 --- a/server/internal/dataexports/repo/queries.sql.go +++ b/server/internal/dataexports/repo/queries.sql.go @@ -118,6 +118,48 @@ func (q *Queries) CreateOtelDestination(ctx context.Context, arg CreateOtelDesti return i, err } +const getActiveOtelRouteDestination = `-- name: GetActiveOtelRouteDestination :one + +SELECT + destination.endpoint_url, + destination.headers_encrypted, + COALESCE(destination.sensitive_data, 'exclude') = 'include' AS include_sensitive_data +FROM data_export_routes AS route +JOIN otel_destinations AS destination + ON destination.organization_id = route.organization_id + AND destination.project_id = route.project_id + AND destination.id = route.otel_destination_id +WHERE route.organization_id = $1 + AND route.project_id = $2 + AND route.data_source = $3 + AND route.enabled IS TRUE + AND route.deleted IS FALSE + AND route.otel_destination_id IS NOT NULL + AND destination.deleted IS FALSE +` + +type GetActiveOtelRouteDestinationParams struct { + OrganizationID string + ProjectID uuid.UUID + DataSource string +} + +type GetActiveOtelRouteDestinationRow struct { + EndpointUrl string + HeadersEncrypted pgtype.Text + IncludeSensitiveData bool +} + +// Every query pins both organization_id and project_id. Resource UUIDs and +// tenant-pinned foreign keys are integrity controls, not authorization bounds. +// Resolve the active OTEL destination for one project data source. +func (q *Queries) GetActiveOtelRouteDestination(ctx context.Context, arg GetActiveOtelRouteDestinationParams) (GetActiveOtelRouteDestinationRow, error) { + row := q.db.QueryRow(ctx, getActiveOtelRouteDestination, arg.OrganizationID, arg.ProjectID, arg.DataSource) + var i GetActiveOtelRouteDestinationRow + err := row.Scan(&i.EndpointUrl, &i.HeadersEncrypted, &i.IncludeSensitiveData) + return i, err +} + const getDataExportRouteForUpdate = `-- name: GetDataExportRouteForUpdate :one SELECT id, organization_id, project_id, data_source, enabled, otel_destination_id, created_at, updated_at, deleted_at, deleted FROM data_export_routes @@ -289,8 +331,6 @@ type ListOtelDestinationsParams struct { ProjectID uuid.UUID } -// Every query pins both organization_id and project_id. Resource UUIDs and -// tenant-pinned foreign keys are integrity controls, not authorization bounds. // List active destinations in stable creation order for the management API. func (q *Queries) ListOtelDestinations(ctx context.Context, arg ListOtelDestinationsParams) ([]OtelDestination, error) { rows, err := q.db.Query(ctx, listOtelDestinations, arg.OrganizationID, arg.ProjectID) diff --git a/server/internal/feature/flags.go b/server/internal/feature/flags.go index 1fa5d95b4ad..0ac0a051047 100644 --- a/server/internal/feature/flags.go +++ b/server/internal/feature/flags.go @@ -118,13 +118,6 @@ const ( // plugins.canaryHooksOrgSlugs), independent of this flag, so a PostHog outage // can't strand it on stale hooks. FlagHooksRollout Flag = "hooks-rollout" - - // FlagOTELLogCustomerRelay controls which organizations have normalized - // OTEL logs relayed to customer-defined destinations. It is targeted by - // PostHog organization group (org slug) and fails closed per organization, - // so one unavailable evaluation never changes another organization's - // delivery. - FlagOTELLogCustomerRelay Flag = "otel-log-customer-relay" ) // Variants of FlagAssistantPlatformMCP. Anything else — no variant, an diff --git a/server/internal/otel/dialect/log_claude_code.go b/server/internal/otel/dialect/log_claude_code.go index 63e93125fa9..ffbbe3c5f43 100644 --- a/server/internal/otel/dialect/log_claude_code.go +++ b/server/internal/otel/dialect/log_claude_code.go @@ -12,7 +12,7 @@ func (ClaudeCodeLog) AppliesTo(record *otelv1.InboundLogRecord) bool { } func (ClaudeCodeLog) InputContent(record *otelv1.InboundLogRecord) (string, genaiconv.InputMessages, error) { - key, value := getOneLogAttr(record, "user_prompt") + key, value := getOneLogAttr(record, claudeCodeUserPromptKey) if key == "" || value == "" { return "", nil, nil } @@ -37,12 +37,12 @@ func (ClaudeCodeLog) SessionID(record *otelv1.InboundLogRecord) (string, string, } func (ClaudeCodeLog) ExternalUserEmail(record *otelv1.InboundLogRecord) (string, string, error) { - key, value := getOneLogAttr(record, "user.email") + key, value := getOneLogAttr(record, userEmailKey) return key, value, nil } func (ClaudeCodeLog) ExternalUserID(record *otelv1.InboundLogRecord) (string, string, error) { - key, value := getOneLogAttr(record, "user.account_id") + key, value := getOneLogAttr(record, vendorUserAccountIDKey) return key, value, nil } diff --git a/server/internal/otel/dialect/log_codex.go b/server/internal/otel/dialect/log_codex.go index 88e0d85c7f5..81c4d513364 100644 --- a/server/internal/otel/dialect/log_codex.go +++ b/server/internal/otel/dialect/log_codex.go @@ -20,7 +20,7 @@ func (CodexLog) AppliesTo(record *otelv1.InboundLogRecord) bool { func (CodexLog) InputContent(record *otelv1.InboundLogRecord) (string, genaiconv.InputMessages, error) { switch record.GetEventName() { case codexUserPromptEvent: - key, prompt := getOneLogAttr(record, "prompt") + key, prompt := getOneLogAttr(record, codexPromptKey) if key == "" || prompt == "" || prompt == codexRedactedUserPrompt { return "", nil, nil } @@ -48,12 +48,12 @@ func (CodexLog) SessionID(record *otelv1.InboundLogRecord) (string, string, erro } func (CodexLog) ExternalUserEmail(record *otelv1.InboundLogRecord) (string, string, error) { - key, value := getOneLogAttr(record, "user.email") + key, value := getOneLogAttr(record, userEmailKey) return key, value, nil } func (CodexLog) ExternalUserID(record *otelv1.InboundLogRecord) (string, string, error) { - key, value := getOneLogAttr(record, "user.account_id") + key, value := getOneLogAttr(record, vendorUserAccountIDKey) return key, value, nil } diff --git a/server/internal/otel/dialect/log_semconv.go b/server/internal/otel/dialect/log_semconv.go index 2d5dc93213d..efc3df4e88d 100644 --- a/server/internal/otel/dialect/log_semconv.go +++ b/server/internal/otel/dialect/log_semconv.go @@ -13,11 +13,11 @@ type SemconvLog struct{} func (SemconvLog) AppliesTo(*otelv1.InboundLogRecord) bool { return true } func (SemconvLog) InputContent(record *otelv1.InboundLogRecord) (string, genaiconv.InputMessages, error) { - return semconvLogContent[genaiconv.InputMessages](record, "gen_ai.input.messages") + return semconvLogContent[genaiconv.InputMessages](record, semconvInputMessagesKey) } func (SemconvLog) OutputContent(record *otelv1.InboundLogRecord) (string, genaiconv.OutputMessages, error) { - return semconvLogContent[genaiconv.OutputMessages](record, "gen_ai.output.messages") + return semconvLogContent[genaiconv.OutputMessages](record, semconvOutputMessagesKey) } func (SemconvLog) SessionID(record *otelv1.InboundLogRecord) (string, string, error) { @@ -26,12 +26,12 @@ func (SemconvLog) SessionID(record *otelv1.InboundLogRecord) (string, string, er } func (SemconvLog) ExternalUserEmail(record *otelv1.InboundLogRecord) (string, string, error) { - key, value := getOneLogAttr(record, "user.email") + key, value := getOneLogAttr(record, userEmailKey) return key, value, nil } func (SemconvLog) ExternalUserID(record *otelv1.InboundLogRecord) (string, string, error) { - key, value := getOneLogAttr(record, "user.id") + key, value := getOneLogAttr(record, semconvUserIDKey) return key, value, nil } diff --git a/server/internal/otel/dialect/metric_claude_code.go b/server/internal/otel/dialect/metric_claude_code.go index b98d92155a5..67100a9952c 100644 --- a/server/internal/otel/dialect/metric_claude_code.go +++ b/server/internal/otel/dialect/metric_claude_code.go @@ -18,12 +18,12 @@ func (ClaudeCodeMetric) SessionID(point MetricDataPoint) (string, string, error) } func (ClaudeCodeMetric) ExternalUserID(point MetricDataPoint) (string, string, error) { - key, value := getOneMetricPointAttr(point, "user.account_id") + key, value := getOneMetricPointAttr(point, vendorUserAccountIDKey) return key, value, nil } func (ClaudeCodeMetric) ExternalUserEmail(point MetricDataPoint) (string, string, error) { - key, value := getOneMetricPointAttr(point, "user.email") + key, value := getOneMetricPointAttr(point, userEmailKey) return key, value, nil } diff --git a/server/internal/otel/dialect/metric_codex.go b/server/internal/otel/dialect/metric_codex.go index b5350dae25f..dfe93b1949e 100644 --- a/server/internal/otel/dialect/metric_codex.go +++ b/server/internal/otel/dialect/metric_codex.go @@ -20,12 +20,12 @@ func (CodexMetric) SessionID(point MetricDataPoint) (string, string, error) { } func (CodexMetric) ExternalUserID(point MetricDataPoint) (string, string, error) { - key, value := getOneMetricPointAttr(point, "user.account_id") + key, value := getOneMetricPointAttr(point, vendorUserAccountIDKey) return key, value, nil } func (CodexMetric) ExternalUserEmail(point MetricDataPoint) (string, string, error) { - key, value := getOneMetricPointAttr(point, "user.email") + key, value := getOneMetricPointAttr(point, userEmailKey) return key, value, nil } diff --git a/server/internal/otel/dialect/metric_semconv.go b/server/internal/otel/dialect/metric_semconv.go index 42827953b77..c8e1b6d00c2 100644 --- a/server/internal/otel/dialect/metric_semconv.go +++ b/server/internal/otel/dialect/metric_semconv.go @@ -12,12 +12,12 @@ func (SemconvMetric) SessionID(point MetricDataPoint) (string, string, error) { } func (SemconvMetric) ExternalUserID(point MetricDataPoint) (string, string, error) { - key, value := getOneMetricPointAttr(point, "user.id") + key, value := getOneMetricPointAttr(point, semconvUserIDKey) return key, value, nil } func (SemconvMetric) ExternalUserEmail(point MetricDataPoint) (string, string, error) { - key, value := getOneMetricPointAttr(point, "user.email") + key, value := getOneMetricPointAttr(point, userEmailKey) return key, value, nil } diff --git a/server/internal/otel/dialect/sensitivity.go b/server/internal/otel/dialect/sensitivity.go new file mode 100644 index 00000000000..059b0b0b059 --- /dev/null +++ b/server/internal/otel/dialect/sensitivity.go @@ -0,0 +1,50 @@ +package dialect + +import "strings" + +// Sensitive attribute keys are shared by dialect extraction and relay redaction. +// Dialect keys stay in the exact set even when a prefix also covers them. +const ( + claudeCodeUserPromptKey = "user_prompt" + codexPromptKey = "prompt" + semconvInputMessagesKey = "gen_ai.input.messages" + semconvOutputMessagesKey = "gen_ai.output.messages" + semconvUserIDKey = "user.id" + userEmailKey = "user.email" + vendorUserAccountIDKey = "user.account_id" +) + +var sensitiveDataExactKeys = map[string]struct{}{ + claudeCodeUserPromptKey: {}, + codexPromptKey: {}, + semconvInputMessagesKey: {}, + semconvOutputMessagesKey: {}, + semconvUserIDKey: {}, + userEmailKey: {}, + vendorUserAccountIDKey: {}, + "assistant": {}, + "content": {}, + "gen_ai.system_instructions": {}, + "tool.args": {}, + "tool_result": {}, +} + +var sensitiveDataPrefixes = [...]string{ + "gen_ai.input.", + "gen_ai.output.", + "gen_ai.tool.call.", + "enduser.", + "user.", +} + +func IsSensitiveDataKey(key string) bool { + if _, ok := sensitiveDataExactKeys[key]; ok { + return true + } + for _, prefix := range sensitiveDataPrefixes { + if strings.HasPrefix(key, prefix) { + return true + } + } + return false +} diff --git a/server/internal/otel/dialect/sensitivity_test.go b/server/internal/otel/dialect/sensitivity_test.go new file mode 100644 index 00000000000..ece1a48d3e3 --- /dev/null +++ b/server/internal/otel/dialect/sensitivity_test.go @@ -0,0 +1,41 @@ +package dialect + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestIsSensitiveDataKey(t *testing.T) { + t.Parallel() + + tests := []struct { + key string + sensitive bool + }{ + {key: "gen_ai.input.messages", sensitive: true}, + {key: "gen_ai.output.messages", sensitive: true}, + {key: "gen_ai.tool.call.arguments", sensitive: true}, + {key: "gen_ai.tool.call.result", sensitive: true}, + {key: "gen_ai.system_instructions", sensitive: true}, + {key: "user_prompt", sensitive: true}, + {key: "prompt", sensitive: true}, + {key: "tool.args", sensitive: true}, + {key: "tool_result", sensitive: true}, + {key: "content", sensitive: true}, + {key: "assistant", sensitive: true}, + {key: "enduser.id", sensitive: true}, + {key: "enduser.email", sensitive: true}, + {key: "user.id", sensitive: true}, + {key: "user.account_id", sensitive: true}, + {key: "user.email", sensitive: true}, + {key: "gen_ai.input", sensitive: false}, + {key: "gen_ai.system", sensitive: false}, + {key: "username", sensitive: false}, + {key: "model", sensitive: false}, + } + + for _, test := range tests { + require.Equalf(t, test.sensitive, IsSensitiveDataKey(test.key), "key %q", test.key) + } +} diff --git a/server/internal/otel/dialect/span_claude_code.go b/server/internal/otel/dialect/span_claude_code.go index f0e645c69da..c5d96ed5f47 100644 --- a/server/internal/otel/dialect/span_claude_code.go +++ b/server/internal/otel/dialect/span_claude_code.go @@ -12,7 +12,7 @@ func (e ClaudeCodeSpan) AppliesTo(span *otelv1.InboundSpan) bool { } func (e ClaudeCodeSpan) InputContent(span *otelv1.InboundSpan) (key string, val genaiconv.InputMessages, err error) { - k, v := getOneAttr(span, "user_prompt") + k, v := getOneAttr(span, claudeCodeUserPromptKey) if k == "" || v == "" { return "", nil, nil } @@ -41,12 +41,12 @@ func (e ClaudeCodeSpan) SessionID(span *otelv1.InboundSpan) (key string, val str } func (e ClaudeCodeSpan) ExternalUserEmail(span *otelv1.InboundSpan) (key string, val string, err error) { - key, val = getOneAttr(span, "user.email") + key, val = getOneAttr(span, userEmailKey) return key, val, nil } func (e ClaudeCodeSpan) ExternalUserID(span *otelv1.InboundSpan) (key string, val string, err error) { - key, val = getOneAttr(span, "user.account_id") + key, val = getOneAttr(span, vendorUserAccountIDKey) return key, val, nil } diff --git a/server/internal/otel/dialect/span_semconv.go b/server/internal/otel/dialect/span_semconv.go index 2d98fcbbe03..95d5dea3bf0 100644 --- a/server/internal/otel/dialect/span_semconv.go +++ b/server/internal/otel/dialect/span_semconv.go @@ -15,11 +15,11 @@ func (e SemconvSpan) AppliesTo(span *otelv1.InboundSpan) bool { } func (e SemconvSpan) InputContent(span *otelv1.InboundSpan) (string, genaiconv.InputMessages, error) { - return semconvContent[genaiconv.InputMessages](span, "gen_ai.input.messages") + return semconvContent[genaiconv.InputMessages](span, semconvInputMessagesKey) } func (e SemconvSpan) OutputContent(span *otelv1.InboundSpan) (string, genaiconv.OutputMessages, error) { - return semconvContent[genaiconv.OutputMessages](span, "gen_ai.output.messages") + return semconvContent[genaiconv.OutputMessages](span, semconvOutputMessagesKey) } func semconvContent[T any](span *otelv1.InboundSpan, desired string) (string, T, error) { @@ -99,12 +99,12 @@ func (e SemconvSpan) SessionID(span *otelv1.InboundSpan) (key string, val string } func (e SemconvSpan) ExternalUserEmail(span *otelv1.InboundSpan) (key string, val string, err error) { - key, val = getOneAttr(span, "user.email") + key, val = getOneAttr(span, userEmailKey) return key, val, nil } func (e SemconvSpan) ExternalUserID(span *otelv1.InboundSpan) (key string, val string, err error) { - key, val = getOneAttr(span, "user.id") + key, val = getOneAttr(span, semconvUserIDKey) return key, val, nil } diff --git a/server/internal/otel/handler_ilog_transforms_test.go b/server/internal/otel/handler_ilog_transforms_test.go index 62240a2180a..694ff1bf2fc 100644 --- a/server/internal/otel/handler_ilog_transforms_test.go +++ b/server/internal/otel/handler_ilog_transforms_test.go @@ -157,7 +157,7 @@ func TestMaxSizeLogRecordFitsRelayExportAfterFullEnrichment(t *testing.T) { require.Positive(t, attributes[string(TokensCountKey)].GetIntValue()) require.NotEmpty(t, attributes[string(TokensCodecKey)].GetStringValue()) - request, err := newLogRelayExportRequest([]*otelv1.LogRecord{published}) + request, err := newLogRelayExportRequest([]*otelv1.LogRecord{published}, true) require.NoError(t, err) require.LessOrEqual(t, proto.Size(request), maxLogRelayExportBytes) } diff --git a/server/internal/otel/handler_log_relay.go b/server/internal/otel/handler_log_relay.go index b6e25bcf9f1..28e78abaa8e 100644 --- a/server/internal/otel/handler_log_relay.go +++ b/server/internal/otel/handler_log_relay.go @@ -6,8 +6,8 @@ import ( "fmt" "log/slog" + "github.com/google/uuid" "github.com/jackc/pgx/v5/pgxpool" - otelv1 "github.com/speakeasy-api/gram/infra/gen/gram/otel/v1" "go.opentelemetry.io/otel/metric" collectorlogsv1 "go.opentelemetry.io/proto/otlp/collector/logs/v1" @@ -20,7 +20,6 @@ import ( "github.com/speakeasy-api/gram/server/internal/attr" "github.com/speakeasy-api/gram/server/internal/encryption" - "github.com/speakeasy-api/gram/server/internal/feature" "github.com/speakeasy-api/gram/server/internal/guardian" "github.com/speakeasy-api/gram/server/internal/o11y" "github.com/speakeasy-api/gram/server/internal/streams" @@ -36,17 +35,16 @@ const ( ) type LogRelayHandler struct { - logger *slog.Logger - recordsDropped metric.Int64Counter - recordsFailed metric.Int64Counter - relay *signalRelay - organizationGate logRelayOrganizationGate + logger *slog.Logger + recordsDropped metric.Int64Counter + recordsFailed metric.Int64Counter + relay *signalRelay } type logProvenanceKey struct { source string organizationID string - projectID string + projectID uuid.UUID } type logRelayMessage struct { @@ -65,7 +63,6 @@ func NewLogRelayHandler( readReplica *pgxpool.Pool, encryptionClient *encryption.Client, policy *guardian.Policy, - features feature.Provider, ) *LogRelayHandler { logger = logger.With(attr.SlogComponent("log-relay-handler")) meter := meterProvider.Meter("github.com/speakeasy-api/gram/server/internal/otel") @@ -95,11 +92,6 @@ func NewLogRelayHandler( "/v1/logs", "log", ), - organizationGate: newLogRelayOrganizationGate( - logger, - readReplica, - features, - ), } } @@ -130,7 +122,7 @@ func (h *LogRelayHandler) handleBatch(ctx context.Context, messages []logRelayMe destination *relayDestination err error } - destinations := make(map[string]destinationResult) + destinations := make(map[relayRouteKey]destinationResult) type destinationDelivery struct { destination *relayDestination batch rightSizedProtoBatch[logRelayMessage, *collectorlogsv1.ExportLogsServiceRequest] @@ -138,15 +130,14 @@ func (h *LogRelayHandler) handleBatch(ctx context.Context, messages []logRelayMe deliveries := make([]destinationDelivery, 0, len(groups)) for _, provenanceGroup := range groups { - organizationID := provenanceGroup.key.organizationID - if !h.organizationGate.Enabled(ctx, organizationID) { - continue + routeKey := relayRouteKey{ + organizationID: provenanceGroup.key.organizationID, + projectID: provenanceGroup.key.projectID, } - - result, ok := destinations[provenanceGroup.key.organizationID] + result, ok := destinations[routeKey] if !ok { - result.destination, result.err = h.relay.destinationForOrganization(ctx, provenanceGroup.key.organizationID) - destinations[provenanceGroup.key.organizationID] = result + result.destination, result.err = h.relay.destinationForRoute(ctx, routeKey) + destinations[routeKey] = result } if result.err != nil { err := fmt.Errorf("load log relay destination: %w", result.err) @@ -159,6 +150,7 @@ func (h *LogRelayHandler) handleBatch(ctx context.Context, messages []logRelayMe "load log relay destination", attr.SlogError(result.err), attr.SlogOrganizationID(provenanceGroup.key.organizationID), + attr.SlogProjectID(provenanceGroup.key.projectID.String()), ) continue } @@ -167,7 +159,9 @@ func (h *LogRelayHandler) handleBatch(ctx context.Context, messages []logRelayMe continue } - batches, err := rightSizeProtoBatches(provenanceGroup.messages, maxLogRelayExportBytes, buildLogRelayExport) + batches, err := rightSizeProtoBatches(provenanceGroup.messages, maxLogRelayExportBytes, func(messages []logRelayMessage) (*collectorlogsv1.ExportLogsServiceRequest, error) { + return buildLogRelayExport(messages, result.destination.includeSensitiveData) + }) if err != nil { h.recordDroppedLogs(ctx, len(provenanceGroup.messages), relayReasonInvalid) logger.ErrorContext( @@ -175,6 +169,7 @@ func (h *LogRelayHandler) handleBatch(ctx context.Context, messages []logRelayMe "build log relay exports", attr.SlogError(err), attr.SlogOrganizationID(provenanceGroup.key.organizationID), + attr.SlogProjectID(provenanceGroup.key.projectID.String()), ) continue } @@ -212,6 +207,7 @@ func (h *LogRelayHandler) handleBatch(ctx context.Context, messages []logRelayMe "relay otel logs", attr.SlogError(err), attr.SlogOrganizationID(item.destination.organizationID), + attr.SlogProjectID(item.destination.projectID.String()), attr.SlogURLFull(item.destination.endpoint), ) } @@ -256,11 +252,16 @@ func groupLogsByProvenance(messages []logRelayMessage) ([]logProvenanceGroup, in invalid++ continue } + projectID, err := uuid.Parse(provenance.GetProjectId()) + if err != nil { + invalid++ + continue + } key := logProvenanceKey{ source: provenance.GetSource(), organizationID: provenance.GetOrganizationId(), - projectID: provenance.GetProjectId(), + projectID: projectID, } index, ok := indexes[key] if !ok { @@ -277,12 +278,12 @@ func groupLogsByProvenance(messages []logRelayMessage) ([]logProvenanceGroup, in return groups, invalid } -func buildLogRelayExport(messages []logRelayMessage) (*collectorlogsv1.ExportLogsServiceRequest, error) { +func buildLogRelayExport(messages []logRelayMessage, includeSensitiveData bool) (*collectorlogsv1.ExportLogsServiceRequest, error) { records := make([]*otelv1.LogRecord, len(messages)) for i, message := range messages { records[i] = message.record } - return newLogRelayExportRequest(records) + return newLogRelayExportRequest(records, includeSensitiveData) } func removeGramLogFields(record *logsv1.LogRecord) error { @@ -302,7 +303,7 @@ func removeGramLogFields(record *logsv1.LogRecord) error { return nil } -func newLogRelayExportRequest(records []*otelv1.LogRecord) (*collectorlogsv1.ExportLogsServiceRequest, error) { +func newLogRelayExportRequest(records []*otelv1.LogRecord, includeSensitiveData bool) (*collectorlogsv1.ExportLogsServiceRequest, error) { type scopeGroupKey struct { scope string schemaURL string @@ -403,5 +404,8 @@ func newLogRelayExportRequest(records []*otelv1.LogRecord) (*collectorlogsv1.Exp for i, group := range resourceGroups { request.ResourceLogs[i] = group.resourceLogs } + if !includeSensitiveData { + redactSensitiveOTLP(request) + } return request, nil } diff --git a/server/internal/otel/handler_log_relay_test.go b/server/internal/otel/handler_log_relay_test.go index 07d858aa4f7..24c3311da16 100644 --- a/server/internal/otel/handler_log_relay_test.go +++ b/server/internal/otel/handler_log_relay_test.go @@ -1,8 +1,6 @@ package otel import ( - "context" - "errors" "fmt" "io" "net/http" @@ -13,16 +11,17 @@ import ( "testing" "time" + "github.com/google/uuid" + "github.com/jackc/pgx/v5/pgtype" + otelv1 "github.com/speakeasy-api/gram/infra/gen/gram/otel/v1" "go.opentelemetry.io/otel/metric" collectorlogsv1 "go.opentelemetry.io/proto/otlp/collector/logs/v1" "google.golang.org/protobuf/encoding/protowire" "google.golang.org/protobuf/proto" - "github.com/speakeasy-api/gram/server/internal/feature" "github.com/speakeasy-api/gram/server/internal/guardian" "github.com/speakeasy-api/gram/server/internal/testenv" - "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -40,101 +39,6 @@ type logRelayRequestCapture struct { requests []capturedLogRelayRequest } -type logRelayFeatureFlagCall struct { - flag feature.Flag - distinctID string - groups map[string]string -} - -type logRelayFeatureFlagResult struct { - enabled bool - err error -} - -type logRelayTestFeatureProvider struct { - mu sync.Mutex - results map[string]logRelayFeatureFlagResult - defaultResult logRelayFeatureFlagResult - calls []logRelayFeatureFlagCall -} - -func (p *logRelayTestFeatureProvider) IsFlagEnabled(_ context.Context, flag feature.Flag, distinctID string, groups map[string]string) (bool, error) { - p.mu.Lock() - defer p.mu.Unlock() - p.calls = append(p.calls, logRelayFeatureFlagCall{ - flag: flag, - distinctID: distinctID, - groups: groups, - }) - result, ok := p.results[distinctID] - if !ok { - result = p.defaultResult - } - return result.enabled, result.err -} - -func (p *logRelayTestFeatureProvider) IsFlagEnabledLocal(context.Context, feature.Flag, string, map[string]string, map[string]string) (bool, error) { - return false, nil -} - -func (p *logRelayTestFeatureProvider) FlagPayload(context.Context, feature.Flag, string, map[string]string) ([]byte, error) { - return nil, nil -} - -func (p *logRelayTestFeatureProvider) snapshotCalls() []logRelayFeatureFlagCall { - p.mu.Lock() - defer p.mu.Unlock() - return slices.Clone(p.calls) -} - -func (p *logRelayTestFeatureProvider) setResult(distinctID string, result logRelayFeatureFlagResult) { - p.mu.Lock() - defer p.mu.Unlock() - if p.results == nil { - p.results = make(map[string]logRelayFeatureFlagResult) - } - p.results[distinctID] = result -} - -type logRelayOrganizationSlugResolverFunc func(context.Context, string) (string, error) - -func (f logRelayOrganizationSlugResolverFunc) OrganizationSlug(ctx context.Context, organizationID string) (string, error) { - return f(ctx, organizationID) -} - -type logRelayOrganizationSlugResult struct { - slug string - err error -} - -type logRelayTestOrganizationSlugResolver struct { - mu sync.Mutex - results map[string]logRelayOrganizationSlugResult - calls []string - started chan<- struct{} - release <-chan struct{} -} - -func (r *logRelayTestOrganizationSlugResolver) OrganizationSlug(_ context.Context, organizationID string) (string, error) { - r.mu.Lock() - r.calls = append(r.calls, organizationID) - result := r.results[organizationID] - r.mu.Unlock() - if r.started != nil { - r.started <- struct{}{} - } - if r.release != nil { - <-r.release - } - return result.slug, result.err -} - -func (r *logRelayTestOrganizationSlugResolver) snapshotCalls() []string { - r.mu.Lock() - defer r.mu.Unlock() - return slices.Clone(r.calls) -} - func (c *logRelayRequestCapture) handler(w http.ResponseWriter, r *http.Request) { body, err := io.ReadAll(r.Body) request := &collectorlogsv1.ExportLogsServiceRequest{} @@ -161,375 +65,70 @@ func (c *logRelayRequestCapture) snapshot() []capturedLogRelayRequest { return slices.Clone(c.requests) } -func TestLogRelayHandlerGatesMixedBatchByOrganization(t *testing.T) { +func TestGroupLogsByProvenanceRejectsInvalidProjectIDs(t *testing.T) { t.Parallel() - capture := &logRelayRequestCapture{mu: sync.Mutex{}, requests: nil} - server := httptest.NewServer(http.HandlerFunc(capture.handler)) - t.Cleanup(server.Close) - flags := &logRelayTestFeatureProvider{ - mu: sync.Mutex{}, - results: map[string]logRelayFeatureFlagResult{ - testLogOrganizationID: {enabled: true, err: nil}, - testLogOtherOrganizationID: {enabled: false, err: nil}, - }, - defaultResult: logRelayFeatureFlagResult{enabled: false, err: nil}, - calls: nil, - } - organizations := &logRelayTestOrganizationSlugResolver{ - mu: sync.Mutex{}, - results: map[string]logRelayOrganizationSlugResult{ - testLogOrganizationID: {slug: "enabled-org", err: nil}, - testLogOtherOrganizationID: {slug: "disabled-org", err: nil}, - }, - calls: nil, - started: nil, - release: nil, - } - handler := newLogRelayTestHandlerWithFeatures(t, testenv.NewMeterProvider(t), flags, organizations) - cacheLogRelayTestDestination(t, handler, testLogOrganizationID, server.URL, map[string]string{"X-Customer": "enabled"}) - cacheLogRelayTestDestination(t, handler, testLogOtherOrganizationID, server.URL, map[string]string{"X-Customer": "disabled"}) - - messages, failures := logRelayTestMessages( - relayTestLogRecord("enabled-1", testLogOrganizationID, testLogProjectID, 0), - relayTestLogRecord("disabled", testLogOtherOrganizationID, testLogOtherProjectID, 0), - relayTestLogRecord("enabled-2", testLogOrganizationID, testLogOtherProjectID, 0), - ) - require.NoError(t, handler.handleBatch(t.Context(), messages)) - for _, failure := range failures { - require.NoError(t, failure) - } - messages, failures = logRelayTestMessages( - relayTestLogRecord("enabled-next-batch", testLogOrganizationID, testLogProjectID, 0), + messages, _ := logRelayTestMessages( + nil, + (&otelv1.LogRecord_builder{}).Build(), + relayTestLogRecord("empty-org", "", testLogProjectID, 0), + relayTestLogRecord("missing-project", testLogOrganizationID, "", 0), + relayTestLogRecord("malformed-project", testLogOrganizationID, "malformed", 0), + relayTestLogRecord("valid", testLogOrganizationID, testLogProjectID, 0), ) - require.NoError(t, handler.handleBatch(t.Context(), messages)) - require.NoError(t, failures[0]) - requests := capture.snapshot() - require.Len(t, requests, 3) - for _, request := range requests { - require.Equal(t, "enabled", request.customer) - } - require.Equal(t, []logRelayFeatureFlagCall{ - { - flag: feature.FlagOTELLogCustomerRelay, - distinctID: testLogOrganizationID, - groups: feature.OrgProjectGroups("enabled-org", ""), - }, - { - flag: feature.FlagOTELLogCustomerRelay, - distinctID: testLogOtherOrganizationID, - groups: feature.OrgProjectGroups("disabled-org", ""), - }, - }, flags.snapshotCalls()) - require.Equal(t, []string{testLogOrganizationID, testLogOtherOrganizationID}, organizations.snapshotCalls()) + groups, invalid := groupLogsByProvenance(messages) + require.Equal(t, 5, invalid) + require.Len(t, groups, 1) + require.Equal(t, uuid.MustParse(testLogProjectID), groups[0].key.projectID) } -func TestLogRelayHandlerIsolatesFlagEvaluationFailureWithinBatch(t *testing.T) { +func TestLogRelayHandlerFailsOnlyProjectWithMalformedDestinationHeaders(t *testing.T) { t.Parallel() capture := &logRelayRequestCapture{mu: sync.Mutex{}, requests: nil} server := httptest.NewServer(http.HandlerFunc(capture.handler)) t.Cleanup(server.Close) - flags := &logRelayTestFeatureProvider{ - mu: sync.Mutex{}, - results: map[string]logRelayFeatureFlagResult{ - testLogOrganizationID: {enabled: true, err: nil}, - testLogOtherOrganizationID: {enabled: false, err: errors.New("evaluate flag")}, - }, - defaultResult: logRelayFeatureFlagResult{enabled: false, err: nil}, - calls: nil, - } - organizations := &logRelayTestOrganizationSlugResolver{ - mu: sync.Mutex{}, - results: map[string]logRelayOrganizationSlugResult{ - testLogOrganizationID: {slug: "enabled-org", err: nil}, - testLogOtherOrganizationID: {slug: "error-org", err: nil}, - testLogThirdOrganizationID: {slug: "", err: errors.New("resolve organization")}, - }, - calls: nil, - started: nil, - release: nil, - } - handler := newLogRelayTestHandlerWithFeatures(t, testenv.NewMeterProvider(t), flags, organizations) - cacheLogRelayTestDestination(t, handler, testLogOrganizationID, server.URL, map[string]string{"X-Customer": "enabled"}) - cacheLogRelayTestDestination(t, handler, testLogOtherOrganizationID, server.URL, map[string]string{"X-Customer": "error"}) - cacheLogRelayTestDestination(t, handler, testLogThirdOrganizationID, server.URL, map[string]string{"X-Customer": "missing"}) - messages, failures := logRelayTestMessages( - relayTestLogRecord("enabled", testLogOrganizationID, testLogProjectID, 0), - relayTestLogRecord("evaluation-error", testLogOtherOrganizationID, testLogOtherProjectID, 0), - relayTestLogRecord("missing-flag", testLogThirdOrganizationID, testLogThirdProjectID, 0), - ) - require.NoError(t, handler.handleBatch(t.Context(), messages)) - for _, failure := range failures { - require.NoError(t, failure) - } - messages, failures = logRelayTestMessages( - relayTestLogRecord("evaluation-error-next-batch", testLogOtherOrganizationID, testLogOtherProjectID, 0), - relayTestLogRecord("organization-error-next-batch", testLogThirdOrganizationID, testLogThirdProjectID, 0), - ) - require.NoError(t, handler.handleBatch(t.Context(), messages)) - for _, failure := range failures { - require.NoError(t, failure) - } - - requests := capture.snapshot() - require.Len(t, requests, 1) - require.Equal(t, "enabled", requests[0].customer) - require.Equal(t, []logRelayFeatureFlagCall{ - { - flag: feature.FlagOTELLogCustomerRelay, - distinctID: testLogOrganizationID, - groups: feature.OrgProjectGroups("enabled-org", ""), - }, - { - flag: feature.FlagOTELLogCustomerRelay, - distinctID: testLogOtherOrganizationID, - groups: feature.OrgProjectGroups("error-org", ""), - }, - { - flag: feature.FlagOTELLogCustomerRelay, - distinctID: testLogOtherOrganizationID, - groups: feature.OrgProjectGroups("error-org", ""), - }, - }, flags.snapshotCalls()) - require.Equal( + db, enc, _ := newRelayRouteTest(t, "/v1/logs") + malformedProjectID := createRelayTestProject(t, db, "org-test") + healthyProjectID := createRelayTestProject(t, db, "org-test") + malformedDestination := createRelayTestDestination( t, - []string{ - testLogOrganizationID, - testLogOtherOrganizationID, - testLogThirdOrganizationID, - testLogOtherOrganizationID, - testLogThirdOrganizationID, - }, - organizations.snapshotCalls(), + db, + "org-test", + malformedProjectID, + "https://collector.example.test", + pgtype.Text{String: "not-valid-ciphertext", Valid: true}, + "exclude", ) -} - -func TestLogRelayOrganizationGateCoalescesConcurrentLookups(t *testing.T) { - t.Parallel() - - started := make(chan struct{}) - release := make(chan struct{}) - flags := &logRelayTestFeatureProvider{ - mu: sync.Mutex{}, - results: map[string]logRelayFeatureFlagResult{ - testLogOrganizationID: {enabled: true, err: nil}, - }, - defaultResult: logRelayFeatureFlagResult{enabled: false, err: nil}, - calls: nil, - } - organizations := &logRelayTestOrganizationSlugResolver{ - mu: sync.Mutex{}, - results: map[string]logRelayOrganizationSlugResult{ - testLogOrganizationID: {slug: "enabled-org", err: nil}, - }, - calls: nil, - started: started, - release: release, - } - gate := newCachedLogRelayOrganizationGate( - testenv.NewLogger(t), - flags, - organizations, - logRelayOrganizationGateMaxSize, - logRelayOrganizationGateCacheTTL, - logRelayOrganizationGateLookupTimeout, - ) - - const callers = 32 - begin := make(chan struct{}) - results := make(chan bool, callers) - var ready sync.WaitGroup - ready.Add(callers) - for range callers { - go func() { - ready.Done() - <-begin - results <- gate.Enabled(t.Context(), testLogOrganizationID) - }() - } - ready.Wait() - close(begin) - <-started - close(release) - - for range callers { - require.True(t, <-results) - } - require.Equal(t, []string{testLogOrganizationID}, organizations.snapshotCalls()) - require.Len(t, flags.snapshotCalls(), 1) -} - -func TestLogRelayOrganizationGateCallerCancellationDoesNotCancelSharedLookup(t *testing.T) { - t.Parallel() - - started := make(chan struct{}) - release := make(chan struct{}) - var organizationCalls atomic.Int64 - organizations := logRelayOrganizationSlugResolverFunc(func(ctx context.Context, _ string) (string, error) { - if organizationCalls.Add(1) == 1 { - close(started) - } - select { - case <-release: - return "enabled-org", nil - case <-ctx.Done(): - return "", ctx.Err() - } - }) - flags := &logRelayTestFeatureProvider{ - mu: sync.Mutex{}, - results: map[string]logRelayFeatureFlagResult{ - testLogOrganizationID: {enabled: true, err: nil}, - }, - defaultResult: logRelayFeatureFlagResult{enabled: false, err: nil}, - calls: nil, - } - gate := newCachedLogRelayOrganizationGate( - testenv.NewLogger(t), - flags, - organizations, - logRelayOrganizationGateMaxSize, - logRelayOrganizationGateCacheTTL, - time.Second, - ) - - callerCtx, cancelCaller := context.WithCancel(t.Context()) - firstResult := make(chan bool, 1) - go func() { - firstResult <- gate.Enabled(callerCtx, testLogOrganizationID) - }() - <-started - cancelCaller() - require.False(t, <-firstResult) - - secondResult := make(chan bool, 1) - go func() { - secondResult <- gate.Enabled(t.Context(), testLogOrganizationID) - }() - close(release) - require.True(t, <-secondResult) - require.True(t, gate.Enabled(t.Context(), testLogOrganizationID)) - require.Equal(t, int64(1), organizationCalls.Load()) - require.Len(t, flags.snapshotCalls(), 1) -} - -func TestLogRelayOrganizationGateBoundsLookupAndDoesNotCacheTimeout(t *testing.T) { - t.Parallel() - - finished := make(chan struct{}) - var organizationCalls atomic.Int64 - organizations := logRelayOrganizationSlugResolverFunc(func(ctx context.Context, _ string) (string, error) { - if organizationCalls.Add(1) == 1 { - <-ctx.Done() - close(finished) - return "", ctx.Err() - } - return "enabled-org", nil - }) - flags := &logRelayTestFeatureProvider{ - mu: sync.Mutex{}, - results: map[string]logRelayFeatureFlagResult{ - testLogOrganizationID: {enabled: true, err: nil}, - }, - defaultResult: logRelayFeatureFlagResult{enabled: false, err: nil}, - calls: nil, - } - gate := newCachedLogRelayOrganizationGate( - testenv.NewLogger(t), - flags, - organizations, - logRelayOrganizationGateMaxSize, - logRelayOrganizationGateCacheTTL, - 50*time.Millisecond, - ) - - require.False(t, gate.Enabled(t.Context(), testLogOrganizationID)) - <-finished - require.EventuallyWithT(t, func(collect *assert.CollectT) { - assert.True(collect, gate.Enabled(t.Context(), testLogOrganizationID)) - }, time.Second, 5*time.Millisecond) - require.True(t, gate.Enabled(t.Context(), testLogOrganizationID)) - require.Equal(t, int64(2), organizationCalls.Load()) - require.Len(t, flags.snapshotCalls(), 1) -} - -func TestLogRelayOrganizationGateDoesNotCacheFeatureProviderError(t *testing.T) { - t.Parallel() - - var organizationCalls atomic.Int64 - organizations := logRelayOrganizationSlugResolverFunc(func(context.Context, string) (string, error) { - organizationCalls.Add(1) - return "enabled-org", nil - }) - flags := &logRelayTestFeatureProvider{ - mu: sync.Mutex{}, - results: map[string]logRelayFeatureFlagResult{ - testLogOrganizationID: {enabled: false, err: errors.New("evaluate flag")}, - }, - defaultResult: logRelayFeatureFlagResult{enabled: false, err: nil}, - calls: nil, - } - gate := newCachedLogRelayOrganizationGate( - testenv.NewLogger(t), - flags, - organizations, - logRelayOrganizationGateMaxSize, - logRelayOrganizationGateCacheTTL, - time.Second, + healthyDestination := createRelayTestDestination( + t, + db, + "org-test", + healthyProjectID, + server.URL, + encryptRelayTestHeaders(t, enc, map[string]string{"X-Customer": "healthy"}), + "include", ) + createRelayTestRoute(t, db, "org-test", malformedProjectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: malformedDestination.ID, Valid: true}) + createRelayTestRoute(t, db, "org-test", healthyProjectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: healthyDestination.ID, Valid: true}) - require.False(t, gate.Enabled(t.Context(), testLogOrganizationID)) - flags.setResult(testLogOrganizationID, logRelayFeatureFlagResult{enabled: true, err: nil}) - require.True(t, gate.Enabled(t.Context(), testLogOrganizationID)) - require.True(t, gate.Enabled(t.Context(), testLogOrganizationID)) - require.Equal(t, int64(2), organizationCalls.Load()) - require.Len(t, flags.snapshotCalls(), 2) -} - -func TestLogRelayOrganizationGateBoundsCachedOrganizations(t *testing.T) { - t.Parallel() - - flags := &logRelayTestFeatureProvider{ - mu: sync.Mutex{}, - results: nil, - defaultResult: logRelayFeatureFlagResult{enabled: true, err: nil}, - calls: nil, - } - organizations := &logRelayTestOrganizationSlugResolver{ - mu: sync.Mutex{}, - results: map[string]logRelayOrganizationSlugResult{ - "organization-1": {slug: "org-1", err: nil}, - "organization-2": {slug: "org-2", err: nil}, - "organization-3": {slug: "org-3", err: nil}, - }, - calls: nil, - started: nil, - release: nil, - } - gate := newCachedLogRelayOrganizationGate( - testenv.NewLogger(t), - flags, - organizations, - 2, - time.Hour, - logRelayOrganizationGateLookupTimeout, + policy, err := guardian.NewUnsafePolicy(testenv.NewTracerProvider(t), nil) + require.NoError(t, err) + handler := NewLogRelayHandler(testenv.NewLogger(t), testenv.NewMeterProvider(t), db, enc, policy) + messages, failures := logRelayTestMessages( + relayTestLogRecord("malformed", "org-test", malformedProjectID.String(), 0), + relayTestLogRecord("healthy", "org-test", healthyProjectID.String(), 0), ) - require.True(t, gate.Enabled(t.Context(), "organization-1")) - require.True(t, gate.Enabled(t.Context(), "organization-2")) - require.True(t, gate.Enabled(t.Context(), "organization-3")) - require.True(t, gate.Enabled(t.Context(), "organization-1")) - - require.Equal( - t, - []string{"organization-1", "organization-2", "organization-3", "organization-1"}, - organizations.snapshotCalls(), - ) - require.Len(t, flags.snapshotCalls(), 4) - require.Equal(t, 2, gate.enabled.Len()) + require.NoError(t, handler.handleBatch(t.Context(), messages)) + require.ErrorContains(t, failures[0], "decrypt destination headers") + require.NoError(t, failures[1]) + requests := capture.snapshot() + require.Len(t, requests, 1) + require.Equal(t, "healthy", requests[0].customer) + require.Equal(t, []string{"healthy"}, relayRequestLogBodies(requests[0].request)) } func TestLogRelayHandlerGroupsByProvenanceAndCachesDestinations(t *testing.T) { @@ -539,8 +138,9 @@ func TestLogRelayHandlerGroupsByProvenanceAndCachesDestinations(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(capture.handler)) t.Cleanup(server.Close) handler := newLogRelayTestHandler(t, testenv.NewMeterProvider(t)) - destinationA := cacheLogRelayTestDestination(t, handler, testLogOrganizationID, server.URL, map[string]string{"X-Customer": "a"}) - cacheLogRelayTestDestination(t, handler, testLogOtherOrganizationID, server.URL, map[string]string{"X-Customer": "b"}) + destinationA := cacheLogRelayTestDestination(t, handler, testLogOrganizationID, testLogProjectID, server.URL, map[string]string{"X-Customer": "a"}, true) + cacheLogRelayTestDestination(t, handler, testLogOrganizationID, testLogOtherProjectID, server.URL, map[string]string{"X-Customer": "a"}, true) + cacheLogRelayTestDestination(t, handler, testLogOtherOrganizationID, testLogProjectID, server.URL, map[string]string{"X-Customer": "b"}, true) require.Equal(t, 10*time.Second, destinationA.httpClient.Timeout) messages, failures := logRelayTestMessages( @@ -605,7 +205,7 @@ func TestLogRelayHandlerRightSizesLargeBatchWithoutMixingOrganizations(t *testin expectedNames := make(map[string][]string, len(specs)) records := make([]*otelv1.LogRecord, 0, 16) for _, spec := range specs { - cacheLogRelayTestDestination(t, handler, spec.organizationID, server.URL, map[string]string{"X-Customer": spec.customer}) + cacheLogRelayTestDestination(t, handler, spec.organizationID, spec.projectID, server.URL, map[string]string{"X-Customer": spec.customer}, true) for index := range spec.recordCount { name := fmt.Sprintf("%s-%d", spec.customer, index) record := relayTestLogRecord(name, spec.organizationID, spec.projectID, recordBodyBytes) @@ -655,7 +255,7 @@ func TestLogRelayHandlerLimitsDestinationRequestsToFourMiB(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(capture.handler)) t.Cleanup(server.Close) handler := newLogRelayTestHandler(t, testenv.NewMeterProvider(t)) - cacheLogRelayTestDestination(t, handler, testLogOrganizationID, server.URL, nil) + cacheLogRelayTestDestination(t, handler, testLogOrganizationID, testLogProjectID, server.URL, nil, true) recordBodyBytes := maxOTLPLogRecordBytes / 2 records := []*otelv1.LogRecord{ @@ -665,6 +265,11 @@ func TestLogRelayHandlerLimitsDestinationRequestsToFourMiB(t *testing.T) { relayTestLogRecord("large-4", testLogOrganizationID, testLogProjectID, recordBodyBytes), relayTestLogRecord("large-5", testLogOrganizationID, testLogProjectID, recordBodyBytes), } + for _, record := range records { + record.SetAttributes([]*otelv1.LogRecord_KeyValue{ + relayTestLogAttribute("prompt", "filtered"), + }) + } for _, record := range records { require.LessOrEqual(t, proto.Size(record), maxOTLPLogRecordBytes) } @@ -682,7 +287,12 @@ func TestLogRelayHandlerLimitsDestinationRequestsToFourMiB(t *testing.T) { require.LessOrEqual(t, captured.bodySize, maxLogRelayExportBytes) for _, resourceLogs := range captured.request.GetResourceLogs() { for _, scopeLogs := range resourceLogs.GetScopeLogs() { - deliveredRecords += len(scopeLogs.GetLogRecords()) + for _, record := range scopeLogs.GetLogRecords() { + require.Len(t, record.GetAttributes(), 1) + require.Equal(t, "prompt", record.GetAttributes()[0].GetKey()) + require.Equal(t, "filtered", record.GetAttributes()[0].GetValue().GetStringValue()) + deliveredRecords++ + } } } } @@ -699,10 +309,10 @@ func TestLogRelayDestinationRejectsRequestOverFourMiB(t *testing.T) { })) t.Cleanup(server.Close) handler := newLogRelayTestHandler(t, testenv.NewMeterProvider(t)) - destination := cacheLogRelayTestDestination(t, handler, testLogOrganizationID, server.URL, nil) + destination := cacheLogRelayTestDestination(t, handler, testLogOrganizationID, testLogProjectID, server.URL, nil, true) request, err := newLogRelayExportRequest([]*otelv1.LogRecord{ relayTestLogRecord("oversized", testLogOrganizationID, testLogProjectID, maxLogRelayExportBytes), - }) + }, true) require.NoError(t, err) err = destination.exportWithLimit(t.Context(), request, maxLogRelayExportBytes) @@ -721,7 +331,7 @@ func TestLogRelayHandlerFailsMessagesForRetryableDestination(t *testing.T) { })) t.Cleanup(server.Close) handler := newLogRelayTestHandler(t, testenv.NewMeterProvider(t)) - cacheLogRelayTestDestination(t, handler, testLogOrganizationID, server.URL, nil) + cacheLogRelayTestDestination(t, handler, testLogOrganizationID, testLogProjectID, server.URL, nil, true) messages, failures := logRelayTestMessages( relayTestLogRecord("failed", testLogOrganizationID, testLogProjectID, 0), @@ -741,7 +351,7 @@ func TestNewLogRelayExportRequestDiscardsGramOnlyFields(t *testing.T) { futureGramField = protowire.AppendString(futureGramField, "internal-future-value") record.ProtoReflect().SetUnknown(append(futureOTLPField, futureGramField...)) - request, err := newLogRelayExportRequest([]*otelv1.LogRecord{record}) + request, err := newLogRelayExportRequest([]*otelv1.LogRecord{record}, true) require.NoError(t, err) require.Len(t, request.GetResourceLogs(), 1) require.Len(t, request.GetResourceLogs()[0].GetScopeLogs(), 1) @@ -762,6 +372,52 @@ func TestNewLogRelayExportRequestDiscardsGramOnlyFields(t *testing.T) { } } +func TestLogRelayExportRedactsSensitiveContentWithoutMutatingSource(t *testing.T) { + t.Parallel() + + record := relayTestLogRecord("redacted", testLogOrganizationID, testLogProjectID, 0) + record.SetAttributes([]*otelv1.LogRecord_KeyValue{ + relayTestLogAttribute("gen_ai.input.messages", "input"), + relayTestLogAttribute("gen_ai.output.messages", "output"), + relayTestLogAttribute("user_prompt", "user"), + relayTestLogAttribute("prompt", "prompt"), + relayTestLogAttribute("model", "preserved"), + }) + before := proto.Clone(record) + + request, err := newLogRelayExportRequest([]*otelv1.LogRecord{record}, false) + require.NoError(t, err) + require.True(t, proto.Equal(before, record)) + + converted := request.GetResourceLogs()[0].GetScopeLogs()[0].GetLogRecords()[0] + require.Len(t, converted.GetAttributes(), 5) + for _, attribute := range converted.GetAttributes()[:4] { + require.Equal(t, redactedSensitiveDataValue, attribute.GetValue().GetStringValue()) + } + require.Equal(t, redactedSensitiveDataValue, converted.GetBody().GetStringValue()) + require.Equal(t, "model", converted.GetAttributes()[4].GetKey()) + require.Equal(t, "preserved", converted.GetAttributes()[4].GetValue().GetStringValue()) +} + +func TestLogRelayExportIncludesSensitiveContent(t *testing.T) { + t.Parallel() + + record := relayTestLogRecord("included", testLogOrganizationID, testLogProjectID, 0) + record.SetAttributes([]*otelv1.LogRecord_KeyValue{ + relayTestLogAttribute("gen_ai.input.messages", "input"), + relayTestLogAttribute("gen_ai.output.messages", "output"), + relayTestLogAttribute("user_prompt", "user"), + relayTestLogAttribute("prompt", "prompt"), + relayTestLogAttribute("model", "preserved"), + }) + + request, err := newLogRelayExportRequest([]*otelv1.LogRecord{record}, true) + require.NoError(t, err) + + converted := request.GetResourceLogs()[0].GetScopeLogs()[0].GetLogRecords()[0] + require.Len(t, converted.GetAttributes(), 5) +} + func logRelayTestMessages(records ...*otelv1.LogRecord) ([]logRelayMessage, []error) { failures := make([]error, len(records)) messages := make([]logRelayMessage, len(records)) @@ -778,72 +434,45 @@ func logRelayTestMessages(records ...*otelv1.LogRecord) ([]logRelayMessage, []er } func newLogRelayTestHandler(t *testing.T, meterProvider metric.MeterProvider) *LogRelayHandler { - t.Helper() - features := &logRelayTestFeatureProvider{ - mu: sync.Mutex{}, - results: nil, - defaultResult: logRelayFeatureFlagResult{enabled: true, err: nil}, - calls: nil, - } - organizations := &logRelayTestOrganizationSlugResolver{ - mu: sync.Mutex{}, - results: map[string]logRelayOrganizationSlugResult{ - testLogOrganizationID: {slug: "org-a", err: nil}, - testLogOtherOrganizationID: {slug: "org-b", err: nil}, - testLogThirdOrganizationID: {slug: "org-c", err: nil}, - }, - calls: nil, - started: nil, - release: nil, - } - return newLogRelayTestHandlerWithFeatures(t, meterProvider, features, organizations) -} - -func newLogRelayTestHandlerWithFeatures( - t *testing.T, - meterProvider metric.MeterProvider, - features feature.Provider, - organizations logRelayOrganizationSlugResolver, -) *LogRelayHandler { t.Helper() policy, err := guardian.NewUnsafePolicy(testenv.NewTracerProvider(t), nil) require.NoError(t, err) - handler := NewLogRelayHandler( + return NewLogRelayHandler( testenv.NewLogger(t), meterProvider, nil, testenv.NewEncryptionClient(t), policy, - features, ) - handler.organizationGate = newCachedLogRelayOrganizationGate( - testenv.NewLogger(t), - features, - organizations, - logRelayOrganizationGateMaxSize, - logRelayOrganizationGateCacheTTL, - logRelayOrganizationGateLookupTimeout, - ) - return handler } func cacheLogRelayTestDestination( t *testing.T, handler *LogRelayHandler, organizationID string, + projectID string, baseURL string, headers map[string]string, + includeSensitiveData bool, ) *relayDestination { t.Helper() - destination, err := handler.relay.newDestination(organizationID, baseURL, headers) + key := relayTestRouteKey(organizationID, projectID) + destination, err := handler.relay.newDestination(key, baseURL, headers, includeSensitiveData) require.NoError(t, err) - handler.relay.destinationCache[organizationID] = cachedRelayDestination{ + handler.relay.destinationCache[key] = cachedRelayDestination{ destination: destination, expiresAt: time.Now().Add(time.Hour), } return destination } +func relayTestLogAttribute(key, value string) *otelv1.LogRecord_KeyValue { + return (&otelv1.LogRecord_KeyValue_builder{ + Key: &key, + Value: (&otelv1.LogRecord_AnyValue_builder{StringValue: &value}).Build(), + }).Build() +} + func relayTestLogRecord(body, organizationID, projectID string, bodyBytes int) *otelv1.LogRecord { resourceSchemaURL := "https://opentelemetry.io/schemas/1.27.0" scopeSchemaURL := "https://opentelemetry.io/schemas/1.28.0" diff --git a/server/internal/otel/handler_metric_relay.go b/server/internal/otel/handler_metric_relay.go index 4a3c5f73e34..43bb0bb8e9d 100644 --- a/server/internal/otel/handler_metric_relay.go +++ b/server/internal/otel/handler_metric_relay.go @@ -6,8 +6,8 @@ import ( "fmt" "log/slog" + "github.com/google/uuid" "github.com/jackc/pgx/v5/pgxpool" - otelv1 "github.com/speakeasy-api/gram/infra/gen/gram/otel/v1" "go.opentelemetry.io/otel/metric" collectormetricsv1 "go.opentelemetry.io/proto/otlp/collector/metrics/v1" @@ -43,7 +43,7 @@ type MetricRelayHandler struct { type metricProvenanceKey struct { source string organizationID string - projectID string + projectID uuid.UUID } type metricRelayMessage struct { @@ -120,7 +120,7 @@ func (h *MetricRelayHandler) handleBatch(ctx context.Context, messages []metricR destination *relayDestination err error } - destinations := make(map[string]destinationResult) + destinations := make(map[relayRouteKey]destinationResult) type destinationDelivery struct { destination *relayDestination batch rightSizedProtoBatch[metricRelayMessage, *collectormetricsv1.ExportMetricsServiceRequest] @@ -128,10 +128,14 @@ func (h *MetricRelayHandler) handleBatch(ctx context.Context, messages []metricR deliveries := make([]destinationDelivery, 0, len(groups)) for _, provenanceGroup := range groups { - result, ok := destinations[provenanceGroup.key.organizationID] + routeKey := relayRouteKey{ + organizationID: provenanceGroup.key.organizationID, + projectID: provenanceGroup.key.projectID, + } + result, ok := destinations[routeKey] if !ok { - result.destination, result.err = h.relay.destinationForOrganization(ctx, provenanceGroup.key.organizationID) - destinations[provenanceGroup.key.organizationID] = result + result.destination, result.err = h.relay.destinationForRoute(ctx, routeKey) + destinations[routeKey] = result } if result.err != nil { err := fmt.Errorf("load metric relay destination: %w", result.err) @@ -144,6 +148,7 @@ func (h *MetricRelayHandler) handleBatch(ctx context.Context, messages []metricR "load metric relay destination", attr.SlogError(result.err), attr.SlogOrganizationID(provenanceGroup.key.organizationID), + attr.SlogProjectID(provenanceGroup.key.projectID.String()), ) continue } @@ -152,7 +157,9 @@ func (h *MetricRelayHandler) handleBatch(ctx context.Context, messages []metricR continue } - batches, err := rightSizeProtoBatches(provenanceGroup.messages, maxMetricRelayExportBytes, buildMetricRelayExport) + batches, err := rightSizeProtoBatches(provenanceGroup.messages, maxMetricRelayExportBytes, func(messages []metricRelayMessage) (*collectormetricsv1.ExportMetricsServiceRequest, error) { + return buildMetricRelayExport(messages, result.destination.includeSensitiveData) + }) if err != nil { h.recordDroppedMetrics(ctx, len(provenanceGroup.messages), relayReasonInvalid) h.logger.ErrorContext( @@ -160,6 +167,7 @@ func (h *MetricRelayHandler) handleBatch(ctx context.Context, messages []metricR "build metric relay exports", attr.SlogError(err), attr.SlogOrganizationID(provenanceGroup.key.organizationID), + attr.SlogProjectID(provenanceGroup.key.projectID.String()), ) continue } @@ -197,6 +205,7 @@ func (h *MetricRelayHandler) handleBatch(ctx context.Context, messages []metricR "relay otel metrics", attr.SlogError(err), attr.SlogOrganizationID(item.destination.organizationID), + attr.SlogProjectID(item.destination.projectID.String()), attr.SlogURLFull(item.destination.endpoint), ) } @@ -241,11 +250,16 @@ func groupMetricsByProvenance(messages []metricRelayMessage) ([]metricProvenance invalid++ continue } + projectID, err := uuid.Parse(provenance.GetProjectId()) + if err != nil { + invalid++ + continue + } key := metricProvenanceKey{ source: provenance.GetSource(), organizationID: provenance.GetOrganizationId(), - projectID: provenance.GetProjectId(), + projectID: projectID, } index, ok := indexes[key] if !ok { @@ -262,12 +276,12 @@ func groupMetricsByProvenance(messages []metricRelayMessage) ([]metricProvenance return groups, invalid } -func buildMetricRelayExport(messages []metricRelayMessage) (*collectormetricsv1.ExportMetricsServiceRequest, error) { +func buildMetricRelayExport(messages []metricRelayMessage, includeSensitiveData bool) (*collectormetricsv1.ExportMetricsServiceRequest, error) { metrics := make([]*otelv1.Metric, len(messages)) for i, message := range messages { metrics[i] = message.metric } - return newMetricRelayExportRequest(metrics) + return newMetricRelayExportRequest(metrics, includeSensitiveData) } func removeGramMetricFields(item *metricsv1.Metric) error { @@ -287,7 +301,7 @@ func removeGramMetricFields(item *metricsv1.Metric) error { return nil } -func newMetricRelayExportRequest(items []*otelv1.Metric) (*collectormetricsv1.ExportMetricsServiceRequest, error) { +func newMetricRelayExportRequest(items []*otelv1.Metric, includeSensitiveData bool) (*collectormetricsv1.ExportMetricsServiceRequest, error) { type scopeGroupKey struct { scope string schemaURL string @@ -388,5 +402,8 @@ func newMetricRelayExportRequest(items []*otelv1.Metric) (*collectormetricsv1.Ex for i, group := range resourceGroups { request.ResourceMetrics[i] = group.resourceMetrics } + if !includeSensitiveData { + redactSensitiveOTLP(request) + } return request, nil } diff --git a/server/internal/otel/handler_metric_relay_test.go b/server/internal/otel/handler_metric_relay_test.go index b0d9bedfc0b..899ab504862 100644 --- a/server/internal/otel/handler_metric_relay_test.go +++ b/server/internal/otel/handler_metric_relay_test.go @@ -9,6 +9,8 @@ import ( "testing" "time" + "github.com/google/uuid" + otelv1 "github.com/speakeasy-api/gram/infra/gen/gram/otel/v1" "go.opentelemetry.io/otel/metric" collectormetricsv1 "go.opentelemetry.io/proto/otlp/collector/metrics/v1" @@ -57,14 +59,36 @@ func (c *metricRelayRequestCapture) snapshot() []capturedMetricRelayRequest { return slices.Clone(c.requests) } +func TestGroupMetricsByProvenanceRejectsInvalidProjectIDs(t *testing.T) { + t.Parallel() + + messages, _ := metricRelayTestMessages( + nil, + (&otelv1.Metric_builder{}).Build(), + relayTestMetric("empty-org", "", testMetricProjectID), + relayTestMetric("missing-project", testMetricOrganizationID, ""), + relayTestMetric("malformed-project", testMetricOrganizationID, "malformed"), + relayTestMetric("valid", testMetricOrganizationID, testMetricProjectID), + ) + + groups, invalid := groupMetricsByProvenance(messages) + require.Equal(t, 5, invalid) + require.Len(t, groups, 1) + require.Equal(t, uuid.MustParse(testMetricProjectID), groups[0].key.projectID) +} + func TestMetricRelayHandlerPreservesMetricsWithoutMixingProvenance(t *testing.T) { t.Parallel() - capture := &metricRelayRequestCapture{mu: sync.Mutex{}, requests: nil} - server := httptest.NewServer(http.HandlerFunc(capture.handler)) - t.Cleanup(server.Close) + captureA := &metricRelayRequestCapture{mu: sync.Mutex{}, requests: nil} + serverA := httptest.NewServer(http.HandlerFunc(captureA.handler)) + t.Cleanup(serverA.Close) + captureB := &metricRelayRequestCapture{mu: sync.Mutex{}, requests: nil} + serverB := httptest.NewServer(http.HandlerFunc(captureB.handler)) + t.Cleanup(serverB.Close) handler := newMetricRelayTestHandler(t, testenv.NewMeterProvider(t)) - cacheMetricRelayTestDestination(t, handler, testMetricOrganizationID, server.URL) + cacheMetricRelayTestDestination(t, handler, testMetricOrganizationID, testMetricProjectID, serverA.URL, true) + cacheMetricRelayTestDestination(t, handler, testMetricOrganizationID, testLogOtherProjectID, serverB.URL, true) messages, failures := metricRelayTestMessages( relayTestMetric("requests", testMetricOrganizationID, testMetricProjectID), @@ -76,13 +100,23 @@ func TestMetricRelayHandlerPreservesMetricsWithoutMixingProvenance(t *testing.T) require.NoError(t, failure) } - requests := capture.snapshot() - require.Len(t, requests, 2) - var names []string - for _, captured := range requests { - require.NoError(t, captured.err) - require.Equal(t, "/v1/metrics", captured.path) - require.Equal(t, "application/x-protobuf", captured.contentType) + requestsA := captureA.snapshot() + require.Len(t, requestsA, 1) + require.NoError(t, requestsA[0].err) + require.Equal(t, "/v1/metrics", requestsA[0].path) + require.Equal(t, "application/x-protobuf", requestsA[0].contentType) + namesA := relayRequestMetricNames(requestsA[0].request) + slices.Sort(namesA) + require.Equal(t, []string{"latency", "requests"}, namesA) + + requestsB := captureB.snapshot() + require.Len(t, requestsB, 1) + require.NoError(t, requestsB[0].err) + require.Equal(t, "/v1/metrics", requestsB[0].path) + require.Equal(t, "application/x-protobuf", requestsB[0].contentType) + require.Equal(t, []string{"tokens"}, relayRequestMetricNames(requestsB[0].request)) + + for _, captured := range []capturedMetricRelayRequest{requestsA[0], requestsB[0]} { for _, resourceMetrics := range captured.request.GetResourceMetrics() { require.Equal(t, "https://opentelemetry.io/schemas/1.27.0", resourceMetrics.GetSchemaUrl()) require.Equal(t, "service.name", resourceMetrics.GetResource().GetAttributes()[0].GetKey()) @@ -96,15 +130,12 @@ func TestMetricRelayHandlerPreservesMetricsWithoutMixingProvenance(t *testing.T) require.Equal(t, "scope-value", scopeMetrics.GetScope().GetAttributes()[0].GetValue().GetStringValue()) require.Equal(t, uint32(3), scopeMetrics.GetScope().GetDroppedAttributesCount()) for _, item := range scopeMetrics.GetMetrics() { - names = append(names, item.GetName()) require.Equal(t, "model", item.GetGauge().GetDataPoints()[0].GetAttributes()[0].GetKey()) require.Empty(t, item.ProtoReflect().GetUnknown(), "Gram provenance must not reach the customer destination") } } } } - slices.Sort(names) - require.Equal(t, []string{"latency", "requests", "tokens"}, names) } func TestMetricRelayHandlerLimitsDestinationExportsTo512KiB(t *testing.T) { @@ -114,7 +145,7 @@ func TestMetricRelayHandlerLimitsDestinationExportsTo512KiB(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(capture.handler)) t.Cleanup(server.Close) handler := newMetricRelayTestHandler(t, testenv.NewMeterProvider(t)) - cacheMetricRelayTestDestination(t, handler, testMetricOrganizationID, server.URL) + cacheMetricRelayTestDestination(t, handler, testMetricOrganizationID, testMetricProjectID, server.URL, true) items := []*otelv1.Metric{ relayTestMetric("one", testMetricOrganizationID, testMetricProjectID), @@ -147,6 +178,43 @@ func TestMetricRelayHandlerLimitsDestinationExportsTo512KiB(t *testing.T) { require.LessOrEqual(t, captured.bodySize, maxMetricRelayExportBytes) } } +func TestMetricRelayExportAppliesSensitiveDataPolicyWithoutMutatingSource(t *testing.T) { + t.Parallel() + + item := relayTestMetric("redacted", testMetricOrganizationID, testMetricProjectID) + item.GetGauge().GetDataPoints()[0].SetAttributes([]*otelv1.Metric_KeyValue{ + (&otelv1.Metric_KeyValue_builder{ + Key: new("gen_ai.tool.call.arguments"), + Value: (&otelv1.Metric_AnyValue_builder{ + StringValue: new("sensitive"), + }).Build(), + }).Build(), + (&otelv1.Metric_KeyValue_builder{ + Key: new("model"), + Value: (&otelv1.Metric_AnyValue_builder{ + StringValue: new("preserved"), + }).Build(), + }).Build(), + }) + before := proto.Clone(item) + + excluded, err := newMetricRelayExportRequest([]*otelv1.Metric{item}, false) + require.NoError(t, err) + require.True(t, proto.Equal(before, item)) + excludedAttributes := excluded.GetResourceMetrics()[0].GetScopeMetrics()[0].GetMetrics()[0].GetGauge().GetDataPoints()[0].GetAttributes() + require.Len(t, excludedAttributes, 2) + require.Equal(t, "gen_ai.tool.call.arguments", excludedAttributes[0].GetKey()) + require.Equal(t, redactedSensitiveDataValue, excludedAttributes[0].GetValue().GetStringValue()) + require.Equal(t, "model", excludedAttributes[1].GetKey()) + require.Equal(t, "preserved", excludedAttributes[1].GetValue().GetStringValue()) + + included, err := newMetricRelayExportRequest([]*otelv1.Metric{item}, true) + require.NoError(t, err) + includedAttributes := included.GetResourceMetrics()[0].GetScopeMetrics()[0].GetMetrics()[0].GetGauge().GetDataPoints()[0].GetAttributes() + require.Len(t, includedAttributes, 2) + require.Equal(t, "sensitive", includedAttributes[0].GetValue().GetStringValue()) + require.Equal(t, "preserved", includedAttributes[1].GetValue().GetStringValue()) +} func metricRelayTestMessages(items ...*otelv1.Metric) ([]metricRelayMessage, []error) { messages := make([]metricRelayMessage, len(items)) @@ -176,16 +244,36 @@ func newMetricRelayTestHandler(t *testing.T, meterProvider metric.MeterProvider) ) } -func cacheMetricRelayTestDestination(t *testing.T, handler *MetricRelayHandler, organizationID, baseURL string) { +func cacheMetricRelayTestDestination( + t *testing.T, + handler *MetricRelayHandler, + organizationID string, + projectID string, + baseURL string, + includeSensitiveData bool, +) { t.Helper() - destination, err := handler.relay.newDestination(organizationID, baseURL, nil) + key := relayTestRouteKey(organizationID, projectID) + destination, err := handler.relay.newDestination(key, baseURL, nil, includeSensitiveData) require.NoError(t, err) - handler.relay.destinationCache[organizationID] = cachedRelayDestination{ + handler.relay.destinationCache[key] = cachedRelayDestination{ destination: destination, expiresAt: time.Now().Add(time.Hour), } } +func relayRequestMetricNames(request *collectormetricsv1.ExportMetricsServiceRequest) []string { + var names []string + for _, resourceMetrics := range request.GetResourceMetrics() { + for _, scopeMetrics := range resourceMetrics.GetScopeMetrics() { + for _, item := range scopeMetrics.GetMetrics() { + names = append(names, item.GetName()) + } + } + } + return names +} + func relayTestMetric(name, organizationID, projectID string) *otelv1.Metric { resourceSchemaURL := "https://opentelemetry.io/schemas/1.27.0" scopeSchemaURL := "https://opentelemetry.io/schemas/1.28.0" diff --git a/server/internal/otel/handler_span_relay.go b/server/internal/otel/handler_span_relay.go index 8603f74238b..7ff14a68d4e 100644 --- a/server/internal/otel/handler_span_relay.go +++ b/server/internal/otel/handler_span_relay.go @@ -6,8 +6,8 @@ import ( "fmt" "log/slog" + "github.com/google/uuid" "github.com/jackc/pgx/v5/pgxpool" - otelv1 "github.com/speakeasy-api/gram/infra/gen/gram/otel/v1" "go.opentelemetry.io/otel/metric" collectortracev1 "go.opentelemetry.io/proto/otlp/collector/trace/v1" @@ -44,7 +44,7 @@ type SpanRelayHandler struct { type spanProvenanceKey struct { source string organizationID string - projectID string + projectID uuid.UUID organizationSlug string projectSlug string apiKeyID string @@ -126,7 +126,7 @@ func (h *SpanRelayHandler) handleBatch(ctx context.Context, messages []spanRelay destination *relayDestination err error } - destinations := make(map[string]destinationResult) + destinations := make(map[relayRouteKey]destinationResult) type delivery struct { destination *relayDestination request *collectortracev1.ExportTraceServiceRequest @@ -135,10 +135,14 @@ func (h *SpanRelayHandler) handleBatch(ctx context.Context, messages []spanRelay deliveries := make([]delivery, 0, len(groups)) for _, provenanceGroup := range groups { - result, ok := destinations[provenanceGroup.key.organizationID] + routeKey := relayRouteKey{ + organizationID: provenanceGroup.key.organizationID, + projectID: provenanceGroup.key.projectID, + } + result, ok := destinations[routeKey] if !ok { - result.destination, result.err = h.relay.destinationForOrganization(ctx, provenanceGroup.key.organizationID) - destinations[provenanceGroup.key.organizationID] = result + result.destination, result.err = h.relay.destinationForRoute(ctx, routeKey) + destinations[routeKey] = result } if result.err != nil { err := fmt.Errorf("load span relay destination: %w", result.err) @@ -151,6 +155,7 @@ func (h *SpanRelayHandler) handleBatch(ctx context.Context, messages []spanRelay "load span relay destination", attr.SlogError(result.err), attr.SlogOrganizationID(provenanceGroup.key.organizationID), + attr.SlogProjectID(provenanceGroup.key.projectID.String()), ) continue } @@ -163,7 +168,7 @@ func (h *SpanRelayHandler) handleBatch(ctx context.Context, messages []spanRelay for i, message := range provenanceGroup.messages { spans[i] = message.span } - request, err := newRelayExportRequest(spans) + request, err := newRelayExportRequest(spans, result.destination.includeSensitiveData) if err != nil { h.recordDroppedSpans(ctx, len(provenanceGroup.messages), relayReasonInvalid) continue @@ -201,6 +206,7 @@ func (h *SpanRelayHandler) handleBatch(ctx context.Context, messages []spanRelay "relay otel spans", attr.SlogError(err), attr.SlogOrganizationID(item.destination.organizationID), + attr.SlogProjectID(item.destination.projectID.String()), attr.SlogURLFull(item.destination.endpoint), ) } @@ -245,11 +251,16 @@ func groupSpansByProvenance(messages []spanRelayMessage) ([]spanProvenanceGroup, invalid++ continue } + projectID, err := uuid.Parse(provenance.GetProjectId()) + if err != nil { + invalid++ + continue + } key := spanProvenanceKey{ source: provenance.GetSource(), organizationID: provenance.GetOrganizationId(), - projectID: provenance.GetProjectId(), + projectID: projectID, organizationSlug: provenance.GetOrganizationSlug(), projectSlug: provenance.GetProjectSlug(), apiKeyID: provenance.GetApiKeyId(), @@ -287,7 +298,7 @@ func removeGramSpanFields(span *tracev1.Span) error { return nil } -func newRelayExportRequest(spans []*otelv1.Span) (*collectortracev1.ExportTraceServiceRequest, error) { +func newRelayExportRequest(spans []*otelv1.Span, includeSensitiveData bool) (*collectortracev1.ExportTraceServiceRequest, error) { type scopeGroupKey struct { scope string schemaURL string @@ -390,5 +401,8 @@ func newRelayExportRequest(spans []*otelv1.Span) (*collectortracev1.ExportTraceS for i, group := range resourceGroups { request.ResourceSpans[i] = group.resourceSpans } + if !includeSensitiveData { + redactSensitiveOTLP(request) + } return request, nil } diff --git a/server/internal/otel/handler_span_relay_test.go b/server/internal/otel/handler_span_relay_test.go index a232e4af065..7e896773b80 100644 --- a/server/internal/otel/handler_span_relay_test.go +++ b/server/internal/otel/handler_span_relay_test.go @@ -70,15 +70,16 @@ func TestSpanRelayHandlerGroupsByProvenanceAndCachesDestinations(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(capture.handler)) t.Cleanup(server.Close) handler := newRelayTestHandler(t, testenv.NewMeterProvider(t)) - destinationA := cacheRelayTestDestination(t, handler, "org-a", server.URL, map[string]string{"X-Customer": "a"}) - cacheRelayTestDestination(t, handler, "org-b", server.URL, map[string]string{"X-Customer": "b"}) + destinationA := cacheRelayTestDestination(t, handler, "org-a", testLogProjectID, server.URL, map[string]string{"X-Customer": "a"}) + cacheRelayTestDestination(t, handler, "org-a", testLogOtherProjectID, server.URL, map[string]string{"X-Customer": "a"}) + cacheRelayTestDestination(t, handler, "org-b", testLogProjectID, server.URL, map[string]string{"X-Customer": "b"}) require.Equal(t, 10*time.Second, destinationA.httpClient.Timeout) messages, failures := relayTestMessages( - relayTestSpan("a-1", "org-a", "project-1"), - relayTestSpan("a-2", "org-a", "project-1"), - relayTestSpan("a-3", "org-a", "project-2"), - relayTestSpan("b-1", "org-b", "project-1"), + relayTestSpan("a-1", "org-a", testLogProjectID), + relayTestSpan("a-2", "org-a", testLogProjectID), + relayTestSpan("a-3", "org-a", testLogOtherProjectID), + relayTestSpan("b-1", "org-b", testLogProjectID), ) require.NoError(t, handler.handleBatch(t.Context(), messages)) for _, failure := range failures { @@ -87,7 +88,7 @@ func TestSpanRelayHandlerGroupsByProvenanceAndCachesDestinations(t *testing.T) { // A second batch must stay on the in-memory destination cache: this handler // has no database connection, so a miss would fail the test. - messages, failures = relayTestMessages(relayTestSpan("a-4", "org-a", "project-1")) + messages, failures = relayTestMessages(relayTestSpan("a-4", "org-a", testLogProjectID)) require.NoError(t, handler.handleBatch(t.Context(), messages)) require.NoError(t, failures[0]) @@ -132,7 +133,7 @@ func TestNewRelayExportRequestDiscardsGramOnlySpanFields(t *testing.T) { futureGramField = protowire.AppendString(futureGramField, "internal-future-value") span.ProtoReflect().SetUnknown(append(futureOTLPField, futureGramField...)) - request, err := newRelayExportRequest([]*otelv1.Span{span}) + request, err := newRelayExportRequest([]*otelv1.Span{span}, true) require.NoError(t, err) require.Len(t, request.GetResourceSpans(), 1) require.Len(t, request.GetResourceSpans()[0].GetScopeSpans(), 1) @@ -156,6 +157,51 @@ func TestNewRelayExportRequestDiscardsGramOnlySpanFields(t *testing.T) { } } +func TestSpanRelayExportRedactsSensitiveContentWithoutMutatingSource(t *testing.T) { + t.Parallel() + + span := relayTestSpan("redacted", testLogOrganizationID, testLogProjectID) + span.SetAttributes([]*otelv1.Span_KeyValue{ + relayTestSpanAttribute("gen_ai.input.messages", "input"), + relayTestSpanAttribute("gen_ai.output.messages", "output"), + relayTestSpanAttribute("user_prompt", "user"), + relayTestSpanAttribute("prompt", "prompt"), + relayTestSpanAttribute("model", "preserved"), + }) + before := proto.Clone(span) + + request, err := newRelayExportRequest([]*otelv1.Span{span}, false) + require.NoError(t, err) + require.True(t, proto.Equal(before, span)) + + converted := request.GetResourceSpans()[0].GetScopeSpans()[0].GetSpans()[0] + require.Len(t, converted.GetAttributes(), 5) + for _, attribute := range converted.GetAttributes()[:4] { + require.Equal(t, redactedSensitiveDataValue, attribute.GetValue().GetStringValue()) + } + require.Equal(t, "model", converted.GetAttributes()[4].GetKey()) + require.Equal(t, "preserved", converted.GetAttributes()[4].GetValue().GetStringValue()) +} + +func TestSpanRelayExportIncludesSensitiveContent(t *testing.T) { + t.Parallel() + + span := relayTestSpan("included", testLogOrganizationID, testLogProjectID) + span.SetAttributes([]*otelv1.Span_KeyValue{ + relayTestSpanAttribute("gen_ai.input.messages", "input"), + relayTestSpanAttribute("gen_ai.output.messages", "output"), + relayTestSpanAttribute("user_prompt", "user"), + relayTestSpanAttribute("prompt", "prompt"), + relayTestSpanAttribute("model", "preserved"), + }) + + request, err := newRelayExportRequest([]*otelv1.Span{span}, true) + require.NoError(t, err) + + converted := request.GetResourceSpans()[0].GetScopeSpans()[0].GetSpans()[0] + require.Len(t, converted.GetAttributes(), 5) +} + func TestSpanRelayHandlerMetricRecordersIgnoreUnavailableCounters(t *testing.T) { t.Parallel() @@ -180,7 +226,8 @@ func TestSpanRelayHandlerCountsInvalidAndMissingDestinationDrops(t *testing.T) { meterProvider := sdkmetric.NewMeterProvider(sdkmetric.WithReader(reader)) t.Cleanup(func() { require.NoError(t, meterProvider.Shutdown(context.Background())) }) handler := newRelayTestHandler(t, meterProvider) - handler.relay.destinationCache["org-a"] = cachedRelayDestination{ + missingKey := relayTestRouteKey("org-a", testLogProjectID) + handler.relay.destinationCache[missingKey] = cachedRelayDestination{ destination: nil, expiresAt: time.Now().Add(time.Hour), } @@ -188,15 +235,18 @@ func TestSpanRelayHandlerCountsInvalidAndMissingDestinationDrops(t *testing.T) { messages, failures := relayTestMessages( nil, (&otelv1.Span_builder{}).Build(), - relayTestSpan("missing-1", "org-a", "project-1"), - relayTestSpan("missing-2", "org-a", "project-1"), + relayTestSpan("empty-org", "", testLogProjectID), + relayTestSpan("missing-project", "org-a", ""), + relayTestSpan("malformed", "org-a", "malformed"), + relayTestSpan("missing-1", "org-a", testLogProjectID), + relayTestSpan("missing-2", "org-a", testLogProjectID), ) require.NoError(t, handler.handleBatch(t.Context(), messages)) for _, failure := range failures { require.NoError(t, failure) } - require.Equal(t, int64(2), relaySpanCount(t, reader, meterSpanRelaySpansDropped, relayReasonInvalid)) + require.Equal(t, int64(5), relaySpanCount(t, reader, meterSpanRelaySpansDropped, relayReasonInvalid)) require.Equal(t, int64(2), relaySpanCount(t, reader, meterSpanRelaySpansDropped, relayReasonNoDestination)) } @@ -217,12 +267,12 @@ func TestSpanRelayHandlerFailsOnlyMessagesForFailedDestination(t *testing.T) { meterProvider := sdkmetric.NewMeterProvider(sdkmetric.WithReader(reader)) t.Cleanup(func() { require.NoError(t, meterProvider.Shutdown(context.Background())) }) handler := newRelayTestHandler(t, meterProvider) - cacheRelayTestDestination(t, handler, "org-a", failedServer.URL, nil) - cacheRelayTestDestination(t, handler, "org-b", successServer.URL, nil) + cacheRelayTestDestination(t, handler, "org-a", testLogProjectID, failedServer.URL, nil) + cacheRelayTestDestination(t, handler, "org-a", testLogOtherProjectID, successServer.URL, nil) messages, failures := relayTestMessages( - relayTestSpan("failed", "org-a", "project-1"), - relayTestSpan("delivered", "org-b", "project-1"), + relayTestSpan("failed", "org-a", testLogProjectID), + relayTestSpan("delivered", "org-a", testLogOtherProjectID), ) require.NoError(t, handler.handleBatch(t.Context(), messages)) require.ErrorContains(t, failures[0], "503 Service Unavailable") @@ -246,9 +296,9 @@ func TestSpanRelayHandlerDropsPermanentHTTPFailureWithoutRetry(t *testing.T) { meterProvider := sdkmetric.NewMeterProvider(sdkmetric.WithReader(reader)) t.Cleanup(func() { require.NoError(t, meterProvider.Shutdown(context.Background())) }) handler := newRelayTestHandler(t, meterProvider) - cacheRelayTestDestination(t, handler, "org-a", server.URL, nil) + cacheRelayTestDestination(t, handler, "org-a", testLogProjectID, server.URL, nil) - messages, failures := relayTestMessages(relayTestSpan("rejected", "org-a", "project-1")) + messages, failures := relayTestMessages(relayTestSpan("rejected", "org-a", testLogProjectID)) require.NoError(t, handler.handleBatch(t.Context(), messages)) require.NoError(t, failures[0]) require.Equal(t, int64(1), requests.Load()) @@ -273,9 +323,9 @@ func TestSpanRelayHandlerRetriesRateLimitFailure(t *testing.T) { meterProvider := sdkmetric.NewMeterProvider(sdkmetric.WithReader(reader)) t.Cleanup(func() { require.NoError(t, meterProvider.Shutdown(context.Background())) }) handler := newRelayTestHandler(t, meterProvider) - cacheRelayTestDestination(t, handler, "org-a", server.URL, nil) + cacheRelayTestDestination(t, handler, "org-a", testLogProjectID, server.URL, nil) - messages, failures := relayTestMessages(relayTestSpan("rate-limited", "org-a", "project-1")) + messages, failures := relayTestMessages(relayTestSpan("rate-limited", "org-a", testLogProjectID)) require.NoError(t, handler.handleBatch(t.Context(), messages)) require.ErrorContains(t, failures[0], "429 Too Many Requests") require.Equal(t, int64(2), requests.Load()) @@ -317,13 +367,15 @@ func cacheRelayTestDestination( t *testing.T, handler *SpanRelayHandler, organizationID string, + projectID string, baseURL string, headers map[string]string, ) *relayDestination { t.Helper() - destination, err := handler.relay.newDestination(organizationID, baseURL, headers) + key := relayTestRouteKey(organizationID, projectID) + destination, err := handler.relay.newDestination(key, baseURL, headers, true) require.NoError(t, err) - handler.relay.destinationCache[organizationID] = cachedRelayDestination{ + handler.relay.destinationCache[key] = cachedRelayDestination{ destination: destination, expiresAt: time.Now().Add(time.Hour), } @@ -357,6 +409,13 @@ func relaySpanCount( return 0 } +func relayTestSpanAttribute(key, value string) *otelv1.Span_KeyValue { + return (&otelv1.Span_KeyValue_builder{ + Key: &key, + Value: (&otelv1.Span_AnyValue_builder{StringValue: &value}).Build(), + }).Build() +} + func relayTestSpan(name, organizationID, projectID string) *otelv1.Span { resourceSchemaURL := "https://opentelemetry.io/schemas/1.27.0" scopeSchemaURL := "https://opentelemetry.io/schemas/1.28.0" diff --git a/server/internal/otel/impl_metrics_test.go b/server/internal/otel/impl_metrics_test.go index ab81b4c703a..449cff82799 100644 --- a/server/internal/otel/impl_metrics_test.go +++ b/server/internal/otel/impl_metrics_test.go @@ -66,7 +66,7 @@ func TestMetricsPublishesInboundWithoutEnrichment(t *testing.T) { normalized, err := metricFromInbound(published) require.NoError(t, err) - rebuilt, err := newMetricRelayExportRequest([]*otelv1.Metric{normalized}) + rebuilt, err := newMetricRelayExportRequest([]*otelv1.Metric{normalized}, true) require.NoError(t, err) require.True(t, proto.Equal(request, rebuilt), "pipeline must preserve producer metric identity and attributes") } diff --git a/server/internal/otel/log_relay_organization.go b/server/internal/otel/log_relay_organization.go deleted file mode 100644 index 65ec9eb53f2..00000000000 --- a/server/internal/otel/log_relay_organization.go +++ /dev/null @@ -1,186 +0,0 @@ -package otel - -import ( - "context" - "errors" - "fmt" - "log/slog" - "time" - - "github.com/hashicorp/golang-lru/v2/expirable" - "github.com/jackc/pgx/v5" - "golang.org/x/sync/singleflight" - - "github.com/speakeasy-api/gram/server/internal/attr" - "github.com/speakeasy-api/gram/server/internal/database" - "github.com/speakeasy-api/gram/server/internal/feature" - organizationsrepo "github.com/speakeasy-api/gram/server/internal/organizations/repo" -) - -const ( - // A short TTL limits rollout and recovery staleness while still removing - // database and PostHog lookups from the per-batch hot path. - logRelayOrganizationGateCacheTTL = 30 * time.Second - - // The subscriber sees logs from every organization, so cap the long tail of - // inactive organizations without evicting the normal working set. - logRelayOrganizationGateMaxSize = 4096 - - // Keep both shared cache fills and individual callers bounded independently. - logRelayOrganizationGateLookupTimeout = time.Second -) - -type logRelayOrganizationSlugResolver interface { - OrganizationSlug(ctx context.Context, organizationID string) (string, error) -} - -type logRelayOrganizationGate interface { - Enabled(ctx context.Context, organizationID string) bool -} - -type cachedLogRelayOrganizationGate struct { - logger *slog.Logger - features feature.Provider - organizations logRelayOrganizationSlugResolver - enabled *expirable.LRU[string, bool] - loads singleflight.Group - lookupTimeout time.Duration -} - -func newLogRelayOrganizationGate( - logger *slog.Logger, - db database.DBTX, - features feature.Provider, -) *cachedLogRelayOrganizationGate { - return newCachedLogRelayOrganizationGate( - logger, - features, - &postgresLogRelayOrganizationSlugResolver{db: db}, - logRelayOrganizationGateMaxSize, - logRelayOrganizationGateCacheTTL, - logRelayOrganizationGateLookupTimeout, - ) -} - -func newCachedLogRelayOrganizationGate( - logger *slog.Logger, - features feature.Provider, - organizations logRelayOrganizationSlugResolver, - maxSize int, - ttl time.Duration, - lookupTimeout time.Duration, -) *cachedLogRelayOrganizationGate { - return &cachedLogRelayOrganizationGate{ - logger: logger, - features: features, - organizations: organizations, - enabled: expirable.NewLRU[string, bool](maxSize, nil, ttl), - loads: singleflight.Group{}, - lookupTimeout: lookupTimeout, - } -} - -func (g *cachedLogRelayOrganizationGate) Enabled(ctx context.Context, organizationID string) bool { - if organizationID == "" { - return false - } - if enabled, ok := g.enabled.Get(organizationID); ok { - return enabled - } - - results := g.loads.DoChan(organizationID, func() (any, error) { - if enabled, ok := g.enabled.Get(organizationID); ok { - return enabled, nil - } - - lookupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), g.lookupTimeout) - defer cancel() - - enabled, err := g.evaluate(lookupCtx, organizationID) - if err != nil { - return false, err - } - g.enabled.Add(organizationID, enabled) - return enabled, nil - }) - - waitCtx, cancel := context.WithTimeout(ctx, g.lookupTimeout) - defer cancel() - - var result singleflight.Result - select { - case result = <-results: - case <-waitCtx.Done(): - return false - } - if result.Err != nil { - return false - } - - enabled, ok := result.Val.(bool) - if !ok { - g.logger.WarnContext( - ctx, - "resolve customer log relay feature flag", - attr.SlogOrganizationID(organizationID), - attr.SlogError(fmt.Errorf("unexpected result type %T", result.Val)), - ) - return false - } - return enabled -} - -func (g *cachedLogRelayOrganizationGate) evaluate(ctx context.Context, organizationID string) (bool, error) { - organizationSlug, err := g.organizations.OrganizationSlug(ctx, organizationID) - if err != nil { - g.logger.WarnContext( - ctx, - "resolve organization for customer log relay feature flag", - attr.SlogError(err), - attr.SlogOrganizationID(organizationID), - ) - return false, fmt.Errorf("resolve organization slug: %w", err) - } - if organizationSlug == "" { - return false, nil - } - - enabled, err := g.features.IsFlagEnabled( - ctx, - feature.FlagOTELLogCustomerRelay, - organizationID, - feature.OrgProjectGroups(organizationSlug, ""), - ) - if err != nil { - g.logger.WarnContext( - ctx, - "evaluate customer log relay feature flag", - attr.SlogError(err), - attr.SlogOrganizationID(organizationID), - ) - return false, fmt.Errorf("evaluate feature flag: %w", err) - } - return enabled, nil -} - -type postgresLogRelayOrganizationSlugResolver struct { - db database.DBTX -} - -func (r *postgresLogRelayOrganizationSlugResolver) OrganizationSlug( - ctx context.Context, - organizationID string, -) (string, error) { - if organizationID == "" { - return "", nil - } - - organization, err := organizationsrepo.New(r.db).GetOrganizationMetadata(ctx, organizationID) - if errors.Is(err, pgx.ErrNoRows) { - return "", nil - } - if err != nil { - return "", fmt.Errorf("get organization metadata: %w", err) - } - return organization.Slug, nil -} diff --git a/server/internal/otel/sensitive_data.go b/server/internal/otel/sensitive_data.go new file mode 100644 index 00000000000..452d2b2de0c --- /dev/null +++ b/server/internal/otel/sensitive_data.go @@ -0,0 +1,80 @@ +package otel + +import ( + commonv1 "go.opentelemetry.io/proto/otlp/common/v1" + logsv1 "go.opentelemetry.io/proto/otlp/logs/v1" + "google.golang.org/protobuf/proto" + "google.golang.org/protobuf/reflect/protoreflect" + + "github.com/speakeasy-api/gram/server/internal/otel/dialect" +) + +const redactedSensitiveDataValue = "[REDACTED]" + +func redactSensitiveOTLP(message proto.Message) { + if message == nil { + return + } + + switch value := message.(type) { + case *logsv1.LogRecord: + redactUnkeyedLogBody(value.GetBody()) + case *commonv1.KeyValue: + if dialect.IsSensitiveDataKey(value.GetKey()) { + value.Value = &commonv1.AnyValue{ + Value: &commonv1.AnyValue_StringValue{StringValue: redactedSensitiveDataValue}, + } + return + } + } + + reflected := message.ProtoReflect() + if !reflected.IsValid() { + return + } + reflected.Range(func(field protoreflect.FieldDescriptor, value protoreflect.Value) bool { + redactSensitiveField(field, value) + return true + }) +} + +func redactSensitiveField(field protoreflect.FieldDescriptor, value protoreflect.Value) { + switch { + case field.IsMap(): + if isProtobufMessage(field.MapValue().Kind()) { + value.Map().Range(func(_ protoreflect.MapKey, value protoreflect.Value) bool { + redactSensitiveOTLP(value.Message().Interface()) + return true + }) + } + case field.IsList(): + if isProtobufMessage(field.Kind()) { + values := value.List() + for i := range values.Len() { + redactSensitiveOTLP(values.Get(i).Message().Interface()) + } + } + case isProtobufMessage(field.Kind()): + redactSensitiveOTLP(value.Message().Interface()) + } +} + +func isProtobufMessage(kind protoreflect.Kind) bool { + return kind == protoreflect.MessageKind || kind == protoreflect.GroupKind +} + +func redactUnkeyedLogBody(value *commonv1.AnyValue) { + if value == nil { + return + } + switch body := value.GetValue().(type) { + case nil, *commonv1.AnyValue_KvlistValue: + return + case *commonv1.AnyValue_ArrayValue: + for _, item := range body.ArrayValue.GetValues() { + redactUnkeyedLogBody(item) + } + default: + value.Value = &commonv1.AnyValue_StringValue{StringValue: redactedSensitiveDataValue} + } +} diff --git a/server/internal/otel/sensitive_data_test.go b/server/internal/otel/sensitive_data_test.go new file mode 100644 index 00000000000..dc75d7f202a --- /dev/null +++ b/server/internal/otel/sensitive_data_test.go @@ -0,0 +1,249 @@ +package otel + +import ( + "testing" + + collectorlogsv1 "go.opentelemetry.io/proto/otlp/collector/logs/v1" + collectormetricsv1 "go.opentelemetry.io/proto/otlp/collector/metrics/v1" + collectortracev1 "go.opentelemetry.io/proto/otlp/collector/trace/v1" + commonv1 "go.opentelemetry.io/proto/otlp/common/v1" + logsv1 "go.opentelemetry.io/proto/otlp/logs/v1" + metricsv1 "go.opentelemetry.io/proto/otlp/metrics/v1" + resourcev1 "go.opentelemetry.io/proto/otlp/resource/v1" + tracev1 "go.opentelemetry.io/proto/otlp/trace/v1" + + "github.com/stretchr/testify/require" +) + +func TestRedactSensitiveOTLPSpanContainers(t *testing.T) { + t.Parallel() + + resourceAttributes := redactionTestAttributes() + scopeAttributes := redactionTestAttributes() + spanAttributes := redactionTestAttributes() + eventAttributes := redactionTestAttributes() + linkAttributes := redactionTestAttributes() + span := &tracev1.Span{ + Attributes: spanAttributes, + Events: []*tracev1.Span_Event{{Attributes: eventAttributes}}, + Links: []*tracev1.Span_Link{{Attributes: linkAttributes}}, + } + scopeSpans := &tracev1.ScopeSpans{ + Scope: &commonv1.InstrumentationScope{Attributes: scopeAttributes}, + Spans: []*tracev1.Span{span}, + } + request := &collectortracev1.ExportTraceServiceRequest{ + ResourceSpans: []*tracev1.ResourceSpans{{ + Resource: &resourcev1.Resource{Attributes: resourceAttributes}, + ScopeSpans: []*tracev1.ScopeSpans{scopeSpans}, + }}, + } + + redactSensitiveOTLP(request) + + for _, attributes := range [][]*commonv1.KeyValue{ + resourceAttributes, + scopeAttributes, + spanAttributes, + eventAttributes, + linkAttributes, + } { + assertRedactionTestAttributes(t, attributes) + } +} + +func TestRedactSensitiveOTLPLogContainersAndNestedBody(t *testing.T) { + t.Parallel() + + resourceAttributes := redactionTestAttributes() + scopeAttributes := redactionTestAttributes() + logAttributes := redactionTestAttributes() + bodyContent := redactionTestAttribute("content", &commonv1.AnyValue{ + Value: &commonv1.AnyValue_BytesValue{BytesValue: []byte("body secret")}, + }) + bodyAssistant := redactionTestAttribute("assistant", &commonv1.AnyValue{ + Value: &commonv1.AnyValue_StringValue{StringValue: "nested secret"}, + }) + bodySafe := redactionTestAttribute("model", &commonv1.AnyValue{ + Value: &commonv1.AnyValue_StringValue{StringValue: "safe-model"}, + }) + nestedObject := &commonv1.AnyValue{ + Value: &commonv1.AnyValue_KvlistValue{ + KvlistValue: &commonv1.KeyValueList{Values: []*commonv1.KeyValue{bodyAssistant, bodySafe}}, + }, + } + bodyArray := redactionTestAttribute("response", &commonv1.AnyValue{ + Value: &commonv1.AnyValue_ArrayValue{ + ArrayValue: &commonv1.ArrayValue{Values: []*commonv1.AnyValue{nestedObject}}, + }, + }) + body := &commonv1.AnyValue{ + Value: &commonv1.AnyValue_KvlistValue{ + KvlistValue: &commonv1.KeyValueList{Values: []*commonv1.KeyValue{bodyContent, bodyArray}}, + }, + } + record := &logsv1.LogRecord{Attributes: logAttributes, Body: body} + scopeLogs := &logsv1.ScopeLogs{ + Scope: &commonv1.InstrumentationScope{Attributes: scopeAttributes}, + LogRecords: []*logsv1.LogRecord{record}, + } + request := &collectorlogsv1.ExportLogsServiceRequest{ + ResourceLogs: []*logsv1.ResourceLogs{{ + Resource: &resourcev1.Resource{Attributes: resourceAttributes}, + ScopeLogs: []*logsv1.ScopeLogs{scopeLogs}, + }}, + } + + redactSensitiveOTLP(request) + + assertRedactionTestAttributes(t, resourceAttributes) + assertRedactionTestAttributes(t, scopeAttributes) + assertRedactionTestAttributes(t, logAttributes) + require.Equal(t, "content", bodyContent.GetKey()) + require.Equal(t, redactedSensitiveDataValue, bodyContent.GetValue().GetStringValue()) + require.Equal(t, "assistant", bodyAssistant.GetKey()) + require.Equal(t, redactedSensitiveDataValue, bodyAssistant.GetValue().GetStringValue()) + require.Equal(t, "model", bodySafe.GetKey()) + require.Equal(t, "safe-model", bodySafe.GetValue().GetStringValue()) +} +func TestRedactSensitiveOTLPScalarLogBodies(t *testing.T) { + t.Parallel() + + scalarBody := &commonv1.AnyValue{ + Value: &commonv1.AnyValue_StringValue{StringValue: "direct secret"}, + } + arrayScalar := &commonv1.AnyValue{ + Value: &commonv1.AnyValue_BytesValue{BytesValue: []byte("array secret")}, + } + arrayBody := &commonv1.AnyValue{ + Value: &commonv1.AnyValue_ArrayValue{ + ArrayValue: &commonv1.ArrayValue{Values: []*commonv1.AnyValue{arrayScalar}}, + }, + } + request := &collectorlogsv1.ExportLogsServiceRequest{ + ResourceLogs: []*logsv1.ResourceLogs{{ScopeLogs: []*logsv1.ScopeLogs{{LogRecords: []*logsv1.LogRecord{ + {Body: scalarBody}, + {Body: arrayBody}, + }}}}}, + } + + redactSensitiveOTLP(request) + + require.Equal(t, redactedSensitiveDataValue, scalarBody.GetStringValue()) + require.Equal(t, redactedSensitiveDataValue, arrayScalar.GetStringValue()) +} + +func TestRedactSensitiveOTLPMetricContainers(t *testing.T) { + t.Parallel() + + resourceAttributes := redactionTestAttributes() + scopeAttributes := redactionTestAttributes() + metadata := redactionTestAttributes() + gaugePointAttributes := redactionTestAttributes() + gaugeExemplarAttributes := redactionTestAttributes() + sumPointAttributes := redactionTestAttributes() + sumExemplarAttributes := redactionTestAttributes() + histogramPointAttributes := redactionTestAttributes() + histogramExemplarAttributes := redactionTestAttributes() + exponentialPointAttributes := redactionTestAttributes() + exponentialExemplarAttributes := redactionTestAttributes() + summaryPointAttributes := redactionTestAttributes() + metrics := []*metricsv1.Metric{ + { + Metadata: metadata, + Data: &metricsv1.Metric_Gauge{Gauge: &metricsv1.Gauge{ + DataPoints: []*metricsv1.NumberDataPoint{{ + Attributes: gaugePointAttributes, + Exemplars: []*metricsv1.Exemplar{{FilteredAttributes: gaugeExemplarAttributes}}, + }}, + }}, + }, + { + Data: &metricsv1.Metric_Sum{Sum: &metricsv1.Sum{ + DataPoints: []*metricsv1.NumberDataPoint{{ + Attributes: sumPointAttributes, + Exemplars: []*metricsv1.Exemplar{{FilteredAttributes: sumExemplarAttributes}}, + }}, + }}, + }, + { + Data: &metricsv1.Metric_Histogram{Histogram: &metricsv1.Histogram{ + DataPoints: []*metricsv1.HistogramDataPoint{{ + Attributes: histogramPointAttributes, + Exemplars: []*metricsv1.Exemplar{{FilteredAttributes: histogramExemplarAttributes}}, + }}, + }}, + }, + { + Data: &metricsv1.Metric_ExponentialHistogram{ExponentialHistogram: &metricsv1.ExponentialHistogram{ + DataPoints: []*metricsv1.ExponentialHistogramDataPoint{{ + Attributes: exponentialPointAttributes, + Exemplars: []*metricsv1.Exemplar{{FilteredAttributes: exponentialExemplarAttributes}}, + }}, + }}, + }, + { + Data: &metricsv1.Metric_Summary{Summary: &metricsv1.Summary{ + DataPoints: []*metricsv1.SummaryDataPoint{{Attributes: summaryPointAttributes}}, + }}, + }, + } + scopeMetrics := &metricsv1.ScopeMetrics{ + Scope: &commonv1.InstrumentationScope{Attributes: scopeAttributes}, + Metrics: metrics, + } + request := &collectormetricsv1.ExportMetricsServiceRequest{ + ResourceMetrics: []*metricsv1.ResourceMetrics{{ + Resource: &resourcev1.Resource{Attributes: resourceAttributes}, + ScopeMetrics: []*metricsv1.ScopeMetrics{scopeMetrics}, + }}, + } + + redactSensitiveOTLP(request) + + for _, attributes := range [][]*commonv1.KeyValue{ + resourceAttributes, + scopeAttributes, + metadata, + gaugePointAttributes, + gaugeExemplarAttributes, + sumPointAttributes, + sumExemplarAttributes, + histogramPointAttributes, + histogramExemplarAttributes, + exponentialPointAttributes, + exponentialExemplarAttributes, + summaryPointAttributes, + } { + assertRedactionTestAttributes(t, attributes) + } +} + +func redactionTestAttributes() []*commonv1.KeyValue { + return []*commonv1.KeyValue{ + redactionTestAttribute("gen_ai.input.messages", &commonv1.AnyValue{ + Value: &commonv1.AnyValue_StringValue{StringValue: "content secret"}, + }), + redactionTestAttribute("user.email", &commonv1.AnyValue{ + Value: &commonv1.AnyValue_StringValue{StringValue: "identity secret"}, + }), + redactionTestAttribute("model", &commonv1.AnyValue{ + Value: &commonv1.AnyValue_StringValue{StringValue: "safe-model"}, + }), + } +} + +func redactionTestAttribute(key string, value *commonv1.AnyValue) *commonv1.KeyValue { + return &commonv1.KeyValue{Key: key, Value: value} +} + +func assertRedactionTestAttributes(t *testing.T, attributes []*commonv1.KeyValue) { + t.Helper() + require.Len(t, attributes, 3) + require.Equal(t, "gen_ai.input.messages", attributes[0].GetKey()) + require.Equal(t, redactedSensitiveDataValue, attributes[0].GetValue().GetStringValue()) + require.Equal(t, "user.email", attributes[1].GetKey()) + require.Equal(t, redactedSensitiveDataValue, attributes[1].GetValue().GetStringValue()) + require.Equal(t, "model", attributes[2].GetKey()) + require.Equal(t, "safe-model", attributes[2].GetValue().GetStringValue()) +} diff --git a/server/internal/otel/signal_relay.go b/server/internal/otel/signal_relay.go index 9904ed427b7..2f81d07bdc0 100644 --- a/server/internal/otel/signal_relay.go +++ b/server/internal/otel/signal_relay.go @@ -12,21 +12,23 @@ import ( "sync" "time" + "github.com/google/uuid" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" "golang.org/x/sync/singleflight" "google.golang.org/protobuf/proto" + dataexportsrepo "github.com/speakeasy-api/gram/server/internal/dataexports/repo" "github.com/speakeasy-api/gram/server/internal/encryption" "github.com/speakeasy-api/gram/server/internal/guardian" "github.com/speakeasy-api/gram/server/internal/o11y" - forwardingrepo "github.com/speakeasy-api/gram/server/internal/otelforwarding/repo" ) const ( - relayDestinationCacheTTL = 60 * time.Second - relayDestinationCacheMaxEntries = 1024 - maxRelayErrorBodyBytes = 4 * 1024 + relayNonSensitiveDestinationCacheTTL = 60 * time.Second + relayDestinationCacheMaxEntries = 1024 + maxRelayErrorBodyBytes = 4 * 1024 + relayDataSourceProductTelemetry = "product_telemetry" ) type relayReason string @@ -52,23 +54,30 @@ type signalRelay struct { now func() time.Time cacheMu sync.RWMutex - destinationCache map[string]cachedRelayDestination + destinationCache map[relayRouteKey]cachedRelayDestination destinationLoads singleflight.Group } +type relayRouteKey struct { + organizationID string + projectID uuid.UUID +} + type cachedRelayDestination struct { destination *relayDestination expiresAt time.Time } -// relayDestination sends protobuf exports to one organization's configured -// OTLP endpoint using its forwarding headers and HTTP policy. +// relayDestination sends protobuf exports to one project's configured OTLP +// endpoint using its destination headers and HTTP policy. type relayDestination struct { - organizationID string - endpoint string - headers http.Header - httpClient *guardian.HTTPClient - signalName string + organizationID string + projectID uuid.UUID + endpoint string + headers http.Header + httpClient *guardian.HTTPClient + signalName string + includeSensitiveData bool } type relayExportError struct { @@ -100,33 +109,35 @@ func newSignalRelay( signalName: signalName, now: time.Now, cacheMu: sync.RWMutex{}, - destinationCache: make(map[string]cachedRelayDestination), + destinationCache: make(map[relayRouteKey]cachedRelayDestination), destinationLoads: singleflight.Group{}, } } -func (r *signalRelay) destinationForOrganization(ctx context.Context, organizationID string) (*relayDestination, error) { +func (r *signalRelay) destinationForRoute(ctx context.Context, key relayRouteKey) (*relayDestination, error) { now := r.now() - if cached, ok := r.cachedDestination(organizationID, now); ok { + if cached, ok := r.cachedDestination(key, now); ok { return cached, nil } - value, err, _ := r.destinationLoads.Do(organizationID, func() (any, error) { + value, err, _ := r.destinationLoads.Do(relayRouteLoadKey(key), func() (any, error) { now := r.now() - if cached, ok := r.cachedDestination(organizationID, now); ok { - return cachedRelayDestination{destination: cached, expiresAt: now.Add(relayDestinationCacheTTL)}, nil + if cached, ok := r.cachedDestination(key, now); ok { + return cachedRelayDestination{destination: cached, expiresAt: time.Time{}}, nil } - destination, err := r.loadDestination(ctx, organizationID) + destination, err := r.loadDestination(ctx, key) if err != nil { return nil, err } - cached := cachedRelayDestination{ - destination: destination, - expiresAt: now.Add(relayDestinationCacheTTL), + cached := cachedRelayDestination{destination: destination, expiresAt: time.Time{}} + // Only exclude and missing destinations are safe to reuse: stale include + // policy could disclose data after an include-to-exclude change. + if destination == nil || !destination.includeSensitiveData { + cached.expiresAt = now.Add(relayNonSensitiveDestinationCacheTTL) + r.cacheDestination(key, cached, now) } - r.cacheDestination(organizationID, cached, now) return cached, nil }) if err != nil { @@ -140,36 +151,46 @@ func (r *signalRelay) destinationForOrganization(ctx context.Context, organizati return cached.destination, nil } -func (r *signalRelay) loadDestination(ctx context.Context, organizationID string) (*relayDestination, error) { - config, err := forwardingrepo.New(r.readReplica).GetOrgOTELForwardingConfig(ctx, organizationID) +// relayRouteLoadKey converts a structured route key into the string required +// by singleflight.Group. The NUL byte delimits the organization ID from the +// fixed-format project UUID so distinct component pairs cannot concatenate to +// the same key. +func relayRouteLoadKey(key relayRouteKey) string { + return key.organizationID + "\x00" + key.projectID.String() +} + +func (r *signalRelay) loadDestination(ctx context.Context, key relayRouteKey) (*relayDestination, error) { + config, err := dataexportsrepo.New(r.readReplica).GetActiveOtelRouteDestination(ctx, dataexportsrepo.GetActiveOtelRouteDestinationParams{ + OrganizationID: key.organizationID, + ProjectID: key.projectID, + DataSource: relayDataSourceProductTelemetry, + }) if errors.Is(err, pgx.ErrNoRows) { return nil, nil } if err != nil { - return nil, fmt.Errorf("get forwarding settings for organization: %w", err) - } - if config.EndpointUrl == "" || !config.Enabled { - return nil, nil + return nil, fmt.Errorf("get active OTEL route destination: %w", err) } headers := make(map[string]string) if config.HeadersEncrypted.Valid && config.HeadersEncrypted.String != "" { plaintext, err := r.encryptionClient.Decrypt(config.HeadersEncrypted.String) if err != nil { - return nil, fmt.Errorf("decrypt forwarding headers: %w", err) + return nil, fmt.Errorf("decrypt destination headers: %w", err) } if err := json.Unmarshal([]byte(plaintext), &headers); err != nil { - return nil, fmt.Errorf("decode forwarding headers: %w", err) + return nil, fmt.Errorf("decode destination headers: %w", err) } } - return r.newDestination(organizationID, config.EndpointUrl, headers) + return r.newDestination(key, config.EndpointUrl, headers, config.IncludeSensitiveData) } func (r *signalRelay) newDestination( - organizationID string, + key relayRouteKey, baseURL string, headerValues map[string]string, + includeSensitiveData bool, ) (*relayDestination, error) { endpoint, err := url.JoinPath(baseURL, r.endpointPath) if err != nil { @@ -195,54 +216,67 @@ func (r *signalRelay) newDestination( retryConfig.ErrorHandler = func(response *http.Response, err error, _ int) (*http.Response, error) { return response, err } - httpClient := r.policy.PooledClient(guardian.WithRetryConfig(retryConfig)) + // Include destinations are deliberately reloaded for every delivery so an + // include-to-exclude policy change takes effect immediately. Their clients + // are therefore short-lived and must not retain idle connections. Exclude + // destinations live in the cache and use pooled connections across exports. + var httpClient *guardian.HTTPClient + if includeSensitiveData { + httpClient = r.policy.Client(guardian.WithRetryConfig(retryConfig)) + } else { + httpClient = r.policy.PooledClient(guardian.WithRetryConfig(retryConfig)) + } httpClient.Timeout = 10 * time.Second return &relayDestination{ - organizationID: organizationID, - endpoint: endpoint, - headers: headers, - httpClient: httpClient, - signalName: r.signalName, + organizationID: key.organizationID, + projectID: key.projectID, + endpoint: endpoint, + headers: headers, + httpClient: httpClient, + signalName: r.signalName, + includeSensitiveData: includeSensitiveData, }, nil } -func (r *signalRelay) cacheDestination(organizationID string, cached cachedRelayDestination, now time.Time) { +func (r *signalRelay) cacheDestination(key relayRouteKey, cached cachedRelayDestination, now time.Time) { var evicted []*relayDestination r.cacheMu.Lock() - if replaced, ok := r.destinationCache[organizationID]; ok { - delete(r.destinationCache, organizationID) + if replaced, ok := r.destinationCache[key]; ok { + delete(r.destinationCache, key) if replaced.destination != nil && replaced.destination != cached.destination { evicted = append(evicted, replaced.destination) } } - for cachedOrganizationID, candidate := range r.destinationCache { + for cachedKey, candidate := range r.destinationCache { if now.Before(candidate.expiresAt) { continue } - delete(r.destinationCache, cachedOrganizationID) + delete(r.destinationCache, cachedKey) if candidate.destination != nil { evicted = append(evicted, candidate.destination) } } for len(r.destinationCache) >= relayDestinationCacheMaxEntries { - evictOrganizationID := "" + var evictKey relayRouteKey var evictCandidate cachedRelayDestination - for cachedOrganizationID, candidate := range r.destinationCache { - if evictOrganizationID == "" || + evictSet := false + for cachedKey, candidate := range r.destinationCache { + if !evictSet || candidate.expiresAt.Before(evictCandidate.expiresAt) || - (candidate.expiresAt.Equal(evictCandidate.expiresAt) && cachedOrganizationID < evictOrganizationID) { - evictOrganizationID = cachedOrganizationID + (candidate.expiresAt.Equal(evictCandidate.expiresAt) && relayRouteKeyLess(cachedKey, evictKey)) { + evictKey = cachedKey evictCandidate = candidate + evictSet = true } } - delete(r.destinationCache, evictOrganizationID) + delete(r.destinationCache, evictKey) if evictCandidate.destination != nil { evicted = append(evicted, evictCandidate.destination) } } - r.destinationCache[organizationID] = cached + r.destinationCache[key] = cached r.cacheMu.Unlock() for _, destination := range evicted { @@ -250,6 +284,13 @@ func (r *signalRelay) cacheDestination(organizationID string, cached cachedRelay } } +func relayRouteKeyLess(left, right relayRouteKey) bool { + if left.organizationID != right.organizationID { + return left.organizationID < right.organizationID + } + return bytes.Compare(left.projectID[:], right.projectID[:]) < 0 +} + func closeIdleRelayDestination(destination *relayDestination) { if destination == nil || destination.httpClient == nil { return @@ -257,9 +298,9 @@ func closeIdleRelayDestination(destination *relayDestination) { destination.httpClient.CloseIdleConnections() } -func (r *signalRelay) cachedDestination(organizationID string, now time.Time) (*relayDestination, bool) { +func (r *signalRelay) cachedDestination(key relayRouteKey, now time.Time) (*relayDestination, bool) { r.cacheMu.RLock() - cached, ok := r.destinationCache[organizationID] + cached, ok := r.destinationCache[key] if !ok { r.cacheMu.RUnlock() return nil, false @@ -271,9 +312,9 @@ func (r *signalRelay) cachedDestination(organizationID string, now time.Time) (* r.cacheMu.RUnlock() // Recheck under the write lock: another goroutine may have refreshed this - // organization after the expired read above. + // project route after the expired read above. r.cacheMu.Lock() - cached, ok = r.destinationCache[organizationID] + cached, ok = r.destinationCache[key] if !ok { r.cacheMu.Unlock() return nil, false @@ -282,7 +323,7 @@ func (r *signalRelay) cachedDestination(organizationID string, now time.Time) (* r.cacheMu.Unlock() return cached.destination, true } - delete(r.destinationCache, organizationID) + delete(r.destinationCache, key) r.cacheMu.Unlock() closeIdleRelayDestination(cached.destination) diff --git a/server/internal/otel/signal_relay_test.go b/server/internal/otel/signal_relay_test.go index 71ab6e2bce4..e629fa6631c 100644 --- a/server/internal/otel/signal_relay_test.go +++ b/server/internal/otel/signal_relay_test.go @@ -1,6 +1,7 @@ package otel import ( + "encoding/json" "fmt" "io" "net" @@ -11,7 +12,14 @@ import ( "testing" "time" + "github.com/google/uuid" + "github.com/jackc/pgx/v5/pgtype" + "github.com/jackc/pgx/v5/pgxpool" + + dataexportsrepo "github.com/speakeasy-api/gram/server/internal/dataexports/repo" + "github.com/speakeasy-api/gram/server/internal/encryption" "github.com/speakeasy-api/gram/server/internal/guardian" + projectsrepo "github.com/speakeasy-api/gram/server/internal/projects/repo" "github.com/speakeasy-api/gram/server/internal/testenv" "github.com/stretchr/testify/require" collectortracev1 "go.opentelemetry.io/proto/otlp/collector/trace/v1" @@ -36,9 +44,12 @@ func TestSignalRelayDestinationUsesConfiguredSignalEndpoint(t *testing.T) { policy, err := guardian.NewUnsafePolicy(testenv.NewTracerProvider(t), nil) require.NoError(t, err) relay := newSignalRelay(nil, nil, policy, "/v1/metrics", "metric") - destination, err := relay.newDestination("organization-id", server.URL, map[string]string{ - "Authorization": "Bearer test-token", - }) + destination, err := relay.newDestination( + relayTestRouteKey("organization-id", testLogProjectID), + server.URL, + map[string]string{"Authorization": "Bearer test-token"}, + true, + ) require.NoError(t, err) err = destination.export(t.Context(), &collectortracev1.ExportTraceServiceRequest{}) @@ -50,6 +61,248 @@ func TestSignalRelayDestinationUsesConfiguredSignalEndpoint(t *testing.T) { require.Equal(t, "Bearer test-token", requestHeader) } +func TestSignalRelayLoadsActiveDataExportRoute(t *testing.T) { + t.Parallel() + + db, enc, relay := newRelayRouteTest(t, "/v1/traces") + projectID := createRelayTestProject(t, db, "org-test") + headers := encryptRelayTestHeaders(t, enc, map[string]string{"Authorization": "Bearer route"}) + destination := createRelayTestDestination(t, db, "org-test", projectID, "https://collector.example.test/otlp", headers, "include") + createRelayTestRoute(t, db, "org-test", projectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: destination.ID, Valid: true}) + + loaded, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: projectID}) + require.NoError(t, err) + require.NotNil(t, loaded) + require.Equal(t, "https://collector.example.test/otlp/v1/traces", loaded.endpoint) + require.Equal(t, "Bearer route", loaded.headers.Get("Authorization")) + require.True(t, loaded.includeSensitiveData) + require.Equal(t, projectID, loaded.projectID) +} +func TestSignalRelayReloadsRoutesThatMayExportSensitiveData(t *testing.T) { + t.Parallel() + + db, enc, relay := newRelayRouteTest(t, "/v1/traces") + now := time.Date(2026, time.August, 20, 12, 0, 0, 0, time.UTC) + relay.now = func() time.Time { return now } + projectID := createRelayTestProject(t, db, "org-test") + headers := encryptRelayTestHeaders(t, enc, nil) + destination := createRelayTestDestination(t, db, "org-test", projectID, "https://collector.example.test/otlp", headers, "include") + createRelayTestRoute(t, db, "org-test", projectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: destination.ID, Valid: true}) + key := relayRouteKey{organizationID: "org-test", projectID: projectID} + + included, err := relay.destinationForRoute(t.Context(), key) + require.NoError(t, err) + require.True(t, included.includeSensitiveData) + require.NotContains(t, relay.destinationCache, key) + + _, err = dataexportsrepo.New(db).UpdateOtelDestination(t.Context(), dataexportsrepo.UpdateOtelDestinationParams{ + Name: destination.Name, + EndpointUrl: destination.EndpointUrl, + HeadersEncrypted: destination.HeadersEncrypted, + SensitiveData: pgtype.Text{String: "exclude", Valid: true}, + OrganizationID: destination.OrganizationID, + ProjectID: destination.ProjectID, + ID: destination.ID, + }) + require.NoError(t, err) + + excluded, err := relay.destinationForRoute(t.Context(), key) + require.NoError(t, err) + require.False(t, excluded.includeSensitiveData) + require.Contains(t, relay.destinationCache, key) +} + +func TestSignalRelayTreatsUnknownSensitiveDataPolicyAsExclude(t *testing.T) { + t.Parallel() + + db, enc, relay := newRelayRouteTest(t, "/v1/traces") + projectID := createRelayTestProject(t, db, "org-test") + destination := createRelayTestDestination(t, db, "org-test", projectID, "https://collector.example.test", encryptRelayTestHeaders(t, enc, nil), "unknown") + createRelayTestRoute(t, db, "org-test", projectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: destination.ID, Valid: true}) + + loaded, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: projectID}) + require.NoError(t, err) + require.NotNil(t, loaded) + require.False(t, loaded.includeSensitiveData) +} + +func TestSignalRelayRoutesSameOrganizationProjectsIndependently(t *testing.T) { + t.Parallel() + + firstRequests := make(chan string, 1) + firstServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + firstRequests <- request.URL.Path + "|" + request.Header.Get("X-Project") + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(firstServer.Close) + secondRequests := make(chan string, 1) + secondServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + secondRequests <- request.URL.Path + "|" + request.Header.Get("X-Project") + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(secondServer.Close) + + db, enc, relay := newRelayRouteTest(t, "/v1/logs") + firstProjectID := createRelayTestProject(t, db, "org-test") + secondProjectID := createRelayTestProject(t, db, "org-test") + firstDestination := createRelayTestDestination( + t, + db, + "org-test", + firstProjectID, + firstServer.URL, + encryptRelayTestHeaders(t, enc, map[string]string{"X-Project": "first"}), + "exclude", + ) + secondDestination := createRelayTestDestination( + t, + db, + "org-test", + secondProjectID, + secondServer.URL, + encryptRelayTestHeaders(t, enc, map[string]string{"X-Project": "second"}), + "include", + ) + createRelayTestRoute(t, db, "org-test", firstProjectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: firstDestination.ID, Valid: true}) + createRelayTestRoute(t, db, "org-test", secondProjectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: secondDestination.ID, Valid: true}) + + first, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: firstProjectID}) + require.NoError(t, err) + second, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: secondProjectID}) + require.NoError(t, err) + require.False(t, first.includeSensitiveData) + require.True(t, second.includeSensitiveData) + require.NoError(t, first.export(t.Context(), &collectortracev1.ExportTraceServiceRequest{})) + require.NoError(t, second.export(t.Context(), &collectortracev1.ExportTraceServiceRequest{})) + require.Equal(t, "/v1/logs|first", <-firstRequests) + require.Equal(t, "/v1/logs|second", <-secondRequests) +} + +func TestSignalRelayReturnsNoDestinationWithoutRoute(t *testing.T) { + t.Parallel() + + db, _, relay := newRelayRouteTest(t, "/v1/traces") + projectID := createRelayTestProject(t, db, "org-test") + + destination, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: projectID}) + require.NoError(t, err) + require.Nil(t, destination) +} + +func TestSignalRelayReturnsNoDestinationForDisabledRoute(t *testing.T) { + t.Parallel() + + db, enc, relay := newRelayRouteTest(t, "/v1/traces") + projectID := createRelayTestProject(t, db, "org-test") + destination := createRelayTestDestination(t, db, "org-test", projectID, "https://collector.example.test", encryptRelayTestHeaders(t, enc, nil), "exclude") + createRelayTestRoute(t, db, "org-test", projectID, relayDataSourceProductTelemetry, false, uuid.NullUUID{UUID: destination.ID, Valid: true}) + + loaded, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: projectID}) + require.NoError(t, err) + require.Nil(t, loaded) +} + +func TestSignalRelayReturnsNoDestinationWithoutSelectedOtelDestination(t *testing.T) { + t.Parallel() + + db, _, relay := newRelayRouteTest(t, "/v1/traces") + projectID := createRelayTestProject(t, db, "org-test") + createRelayTestRoute(t, db, "org-test", projectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: uuid.Nil, Valid: false}) + + destination, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: projectID}) + require.NoError(t, err) + require.Nil(t, destination) +} + +func TestSignalRelayReturnsNoDestinationForSoftDeletedRoute(t *testing.T) { + t.Parallel() + + db, enc, relay := newRelayRouteTest(t, "/v1/traces") + projectID := createRelayTestProject(t, db, "org-test") + destination := createRelayTestDestination(t, db, "org-test", projectID, "https://collector.example.test", encryptRelayTestHeaders(t, enc, nil), "exclude") + route := createRelayTestRoute(t, db, "org-test", projectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: destination.ID, Valid: true}) + _, err := dataexportsrepo.New(db).SoftDeleteDataExportRoute(t.Context(), dataexportsrepo.SoftDeleteDataExportRouteParams{ + OrganizationID: "org-test", + ProjectID: projectID, + ID: route.ID, + }) + require.NoError(t, err) + + loaded, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: projectID}) + require.NoError(t, err) + require.Nil(t, loaded) +} + +func TestSignalRelayReturnsNoDestinationForSoftDeletedDestination(t *testing.T) { + t.Parallel() + + db, enc, relay := newRelayRouteTest(t, "/v1/traces") + projectID := createRelayTestProject(t, db, "org-test") + destination := createRelayTestDestination(t, db, "org-test", projectID, "https://collector.example.test", encryptRelayTestHeaders(t, enc, nil), "exclude") + createRelayTestRoute(t, db, "org-test", projectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: destination.ID, Valid: true}) + _, err := dataexportsrepo.New(db).SoftDeleteOtelDestination(t.Context(), dataexportsrepo.SoftDeleteOtelDestinationParams{ + OrganizationID: "org-test", + ProjectID: projectID, + ID: destination.ID, + }) + require.NoError(t, err) + + loaded, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: projectID}) + require.NoError(t, err) + require.Nil(t, loaded) +} + +func TestSignalRelayIgnoresRoutesForOtherDataSources(t *testing.T) { + t.Parallel() + + db, enc, relay := newRelayRouteTest(t, "/v1/traces") + projectID := createRelayTestProject(t, db, "org-test") + destination := createRelayTestDestination(t, db, "org-test", projectID, "https://collector.example.test", encryptRelayTestHeaders(t, enc, nil), "exclude") + createRelayTestRoute(t, db, "org-test", projectID, "risk_findings", true, uuid.NullUUID{UUID: destination.ID, Valid: true}) + + loaded, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: projectID}) + require.NoError(t, err) + require.Nil(t, loaded) +} + +func TestSignalRelayCannotResolveAnotherProjectRoute(t *testing.T) { + t.Parallel() + + db, enc, relay := newRelayRouteTest(t, "/v1/traces") + configuredProjectID := createRelayTestProject(t, db, "org-test") + otherProjectID := createRelayTestProject(t, db, "org-test") + destination := createRelayTestDestination(t, db, "org-test", configuredProjectID, "https://collector.example.test", encryptRelayTestHeaders(t, enc, nil), "exclude") + createRelayTestRoute(t, db, "org-test", configuredProjectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: destination.ID, Valid: true}) + + loaded, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: otherProjectID}) + require.NoError(t, err) + require.Nil(t, loaded) + loaded, err = relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "other-org", projectID: configuredProjectID}) + require.NoError(t, err) + require.Nil(t, loaded) +} + +func TestSignalRelayFailsMalformedEncryptedHeaders(t *testing.T) { + t.Parallel() + + db, _, relay := newRelayRouteTest(t, "/v1/traces") + projectID := createRelayTestProject(t, db, "org-test") + destination := createRelayTestDestination( + t, + db, + "org-test", + projectID, + "https://collector.example.test", + pgtype.Text{String: "not-valid-ciphertext", Valid: true}, + "exclude", + ) + createRelayTestRoute(t, db, "org-test", projectID, relayDataSourceProductTelemetry, true, uuid.NullUUID{UUID: destination.ID, Valid: true}) + + loaded, err := relay.destinationForRoute(t.Context(), relayRouteKey{organizationID: "org-test", projectID: projectID}) + require.ErrorContains(t, err, "decrypt destination headers") + require.Nil(t, loaded) +} + func TestSignalRelayDestinationDrainsFailedResponses(t *testing.T) { t.Parallel() @@ -70,7 +323,7 @@ func TestSignalRelayDestinationDrainsFailedResponses(t *testing.T) { policy, err := guardian.NewUnsafePolicy(testenv.NewTracerProvider(t), nil) require.NoError(t, err) relay := newSignalRelay(nil, nil, policy, "/v1/traces", "trace") - destination, err := relay.newDestination("organization-id", server.URL, nil) + destination, err := relay.newDestination(relayTestRouteKey("organization-id", testLogProjectID), server.URL, nil, false) require.NoError(t, err) require.Error(t, destination.export(t.Context(), &collectortracev1.ExportTraceServiceRequest{})) @@ -98,7 +351,7 @@ func TestSignalRelayDestinationSanitizesResponseDiagnostics(t *testing.T) { policy, err := guardian.NewUnsafePolicy(testenv.NewTracerProvider(t), nil) require.NoError(t, err) permanentRelay := newSignalRelay(nil, nil, policy, "/permanent", "trace") - permanentDestination, err := permanentRelay.newDestination("organization-id", server.URL, nil) + permanentDestination, err := permanentRelay.newDestination(relayTestRouteKey("organization-id", testLogProjectID), server.URL, nil, true) require.NoError(t, err) err = permanentDestination.export(t.Context(), &collectortracev1.ExportTraceServiceRequest{}) @@ -108,7 +361,7 @@ func TestSignalRelayDestinationSanitizesResponseDiagnostics(t *testing.T) { require.NotContains(t, err.Error(), "\x07") retryableRelay := newSignalRelay(nil, nil, policy, "/retryable", "trace") - retryableDestination, err := retryableRelay.newDestination("organization-id", server.URL, nil) + retryableDestination, err := retryableRelay.newDestination(relayTestRouteKey("organization-id", testLogProjectID), server.URL, nil, true) require.NoError(t, err) err = retryableDestination.export(t.Context(), &collectortracev1.ExportTraceServiceRequest{}) @@ -123,45 +376,75 @@ func TestSignalRelayCacheReturnsActiveAndRemovesExpiredDestinations(t *testing.T relay := newSignalRelay(nil, nil, nil, "", "trace") active, activeTransport := newTrackedRelayDestination() expired, expiredTransport := newTrackedRelayDestination() - relay.destinationCache["active"] = cachedRelayDestination{ + activeKey := relayTestRouteKey("active", testLogProjectID) + expiredKey := relayTestRouteKey("expired", testLogProjectID) + relay.destinationCache[activeKey] = cachedRelayDestination{ destination: active, expiresAt: now.Add(time.Minute), } - relay.destinationCache["expired"] = cachedRelayDestination{ + relay.destinationCache[expiredKey] = cachedRelayDestination{ destination: expired, expiresAt: now, } - got, ok := relay.cachedDestination("active", now) + got, ok := relay.cachedDestination(activeKey, now) require.True(t, ok) require.Same(t, active, got) require.Zero(t, activeTransport.closeCalls.Load()) - got, ok = relay.cachedDestination("expired", now) + got, ok = relay.cachedDestination(expiredKey, now) require.False(t, ok) require.Nil(t, got) - require.NotContains(t, relay.destinationCache, "expired") + require.NotContains(t, relay.destinationCache, expiredKey) require.Equal(t, int64(1), expiredTransport.closeCalls.Load()) } +func TestSignalRelayCacheIsolatesProjectsWithinOrganization(t *testing.T) { + t.Parallel() + + now := time.Date(2026, time.August, 20, 12, 0, 0, 0, time.UTC) + relay := newSignalRelay(nil, nil, nil, "", "trace") + first, _ := newTrackedRelayDestination() + second, _ := newTrackedRelayDestination() + firstKey := relayTestRouteKey("organization-id", testLogProjectID) + secondKey := relayTestRouteKey("organization-id", testLogOtherProjectID) + relay.destinationCache[firstKey] = cachedRelayDestination{ + destination: first, + expiresAt: now.Add(time.Minute), + } + relay.destinationCache[secondKey] = cachedRelayDestination{ + destination: second, + expiresAt: now.Add(time.Minute), + } + + firstResult, ok := relay.cachedDestination(firstKey, now) + require.True(t, ok) + require.Same(t, first, firstResult) + secondResult, ok := relay.cachedDestination(secondKey, now) + require.True(t, ok) + require.Same(t, second, secondResult) +} + func TestSignalRelayCacheInsertionPrunesExpiredDestinations(t *testing.T) { t.Parallel() now := time.Date(2026, time.August, 20, 12, 0, 0, 0, time.UTC) relay := newSignalRelay(nil, nil, nil, "", "trace") expired, expiredTransport := newTrackedRelayDestination() - relay.destinationCache["expired"] = cachedRelayDestination{ + expiredKey := relayTestRouteKey("expired", testLogProjectID) + newKey := relayTestRouteKey("new", testLogProjectID) + relay.destinationCache[expiredKey] = cachedRelayDestination{ destination: expired, expiresAt: now, } - relay.cacheDestination("new", cachedRelayDestination{ + relay.cacheDestination(newKey, cachedRelayDestination{ destination: nil, expiresAt: now.Add(time.Minute), }, now) - require.NotContains(t, relay.destinationCache, "expired") - require.Contains(t, relay.destinationCache, "new") + require.NotContains(t, relay.destinationCache, expiredKey) + require.Contains(t, relay.destinationCache, newKey) require.Equal(t, int64(1), expiredTransport.closeCalls.Load()) } @@ -172,17 +455,18 @@ func TestSignalRelayCacheReplacementClosesOldDestination(t *testing.T) { relay := newSignalRelay(nil, nil, nil, "", "trace") oldDestination, oldTransport := newTrackedRelayDestination() newDestination, newTransport := newTrackedRelayDestination() - relay.cacheDestination("organization-id", cachedRelayDestination{ + key := relayTestRouteKey("organization-id", testLogProjectID) + relay.cacheDestination(key, cachedRelayDestination{ destination: oldDestination, expiresAt: now.Add(time.Minute), }, now) - relay.cacheDestination("organization-id", cachedRelayDestination{ + relay.cacheDestination(key, cachedRelayDestination{ destination: newDestination, expiresAt: now.Add(2 * time.Minute), }, now) - got, ok := relay.cachedDestination("organization-id", now) + got, ok := relay.cachedDestination(key, now) require.True(t, ok) require.Same(t, newDestination, got) require.Equal(t, int64(1), oldTransport.closeCalls.Load()) @@ -195,25 +479,28 @@ func TestSignalRelayCacheEvictsEarliestExpiryAtCapacity(t *testing.T) { now := time.Date(2026, time.August, 20, 12, 0, 0, 0, time.UTC) relay := newSignalRelay(nil, nil, nil, "", "trace") oldestDestination, oldestTransport := newTrackedRelayDestination() + oldestKey := relayTestRouteKey("organization-0000", testLogProjectID) for i := range relayDestinationCacheMaxEntries { destination := (*relayDestination)(nil) if i == 0 { destination = oldestDestination } - relay.destinationCache[fmt.Sprintf("organization-%04d", i)] = cachedRelayDestination{ + key := relayTestRouteKey(fmt.Sprintf("organization-%04d", i), testLogProjectID) + relay.destinationCache[key] = cachedRelayDestination{ destination: destination, expiresAt: now.Add(time.Duration(i+1) * time.Second), } } - relay.cacheDestination("new-organization", cachedRelayDestination{ + newKey := relayTestRouteKey("new-organization", testLogProjectID) + relay.cacheDestination(newKey, cachedRelayDestination{ destination: nil, expiresAt: now.Add(time.Hour), }, now) require.Len(t, relay.destinationCache, relayDestinationCacheMaxEntries) - require.NotContains(t, relay.destinationCache, "organization-0000") - require.Contains(t, relay.destinationCache, "new-organization") + require.NotContains(t, relay.destinationCache, oldestKey) + require.Contains(t, relay.destinationCache, newKey) require.Equal(t, int64(1), oldestTransport.closeCalls.Load()) } @@ -234,14 +521,99 @@ func newTrackedRelayDestination() (*relayDestination, *closeIdleTrackingRoundTri client := new(http.Client) client.Transport = transport return &relayDestination{ - organizationID: "", - endpoint: "", - headers: nil, - httpClient: client, - signalName: "", + organizationID: "", + projectID: uuid.Nil, + endpoint: "", + headers: nil, + httpClient: client, + signalName: "", + includeSensitiveData: false, }, transport } +func newRelayRouteTest(t *testing.T, endpointPath string) (*pgxpool.Pool, *encryption.Client, *signalRelay) { + t.Helper() + db := newTestDatabase(t) + enc := testenv.NewEncryptionClient(t) + policy, err := guardian.NewUnsafePolicy(testenv.NewTracerProvider(t), nil) + require.NoError(t, err) + return db, enc, newSignalRelay(db, enc, policy, endpointPath, "test") +} + +func createRelayTestProject(t *testing.T, db *pgxpool.Pool, organizationID string) uuid.UUID { + t.Helper() + slug := "relay-" + uuid.NewString()[:8] + project, err := projectsrepo.New(db).CreateProject(t.Context(), projectsrepo.CreateProjectParams{ + Name: "Relay Test", + Slug: slug, + OrganizationID: organizationID, + }) + require.NoError(t, err) + return project.ID +} + +func encryptRelayTestHeaders(t *testing.T, enc *encryption.Client, headers map[string]string) pgtype.Text { + t.Helper() + if headers == nil { + headers = map[string]string{} + } + encoded, err := json.Marshal(headers) + require.NoError(t, err) + ciphertext, err := enc.Encrypt(encoded) + require.NoError(t, err) + return pgtype.Text{String: ciphertext, Valid: true} +} + +func createRelayTestDestination( + t *testing.T, + db *pgxpool.Pool, + organizationID string, + projectID uuid.UUID, + endpointURL string, + headersEncrypted pgtype.Text, + sensitiveData string, +) dataexportsrepo.OtelDestination { + t.Helper() + destination, err := dataexportsrepo.New(db).CreateOtelDestination(t.Context(), dataexportsrepo.CreateOtelDestinationParams{ + OrganizationID: organizationID, + ProjectID: projectID, + Name: "Relay collector", + EndpointUrl: endpointURL, + HeadersEncrypted: headersEncrypted, + SensitiveData: pgtype.Text{String: sensitiveData, Valid: true}, + }) + require.NoError(t, err) + return destination +} + +func createRelayTestRoute( + t *testing.T, + db *pgxpool.Pool, + organizationID string, + projectID uuid.UUID, + dataSource string, + enabled bool, + destinationID uuid.NullUUID, +) dataexportsrepo.DataExportRoute { + t.Helper() + route, err := dataexportsrepo.New(db).CreateDataExportRoute(t.Context(), dataexportsrepo.CreateDataExportRouteParams{ + OrganizationID: organizationID, + ProjectID: projectID, + DataSource: dataSource, + Enabled: enabled, + OtelDestinationID: destinationID, + }) + require.NoError(t, err) + return route +} + +func relayTestRouteKey(organizationID, projectID string) relayRouteKey { + return relayRouteKey{ + organizationID: organizationID, + projectID: uuid.MustParse(projectID), + } +} + func TestClassifyRelayStatusDistinguishesRetryableFailures(t *testing.T) { t.Parallel()