From 5eebaa9ad4ee5eadd82071cecc1b839df9cf1441 Mon Sep 17 00:00:00 2001 From: Zhen Ye Date: Wed, 15 Jul 2026 02:14:37 +0800 Subject: [PATCH] enhance: add WAL trace propagation (#50796) issue: #47404 - message trace context: add trace context serialization and restore/inject helpers for WAL messages and msgstream conversion - WAL append trace: normalize WAL spans for autocommit, txn, broadcast, append, appendimpl, and broadcast callback paths - trace propagation: restore message trace context in producer, broadcast retry, replicate primary/secondary, recovery, flusher, and query/data flowgraph consumers --------- Signed-off-by: chyezh Co-authored-by: Claude Opus 4.6 (1M context) --- CLAUDE.md | 1 + docs/agent_guides/observability/README.md | 24 ++ .../streaming-system/streaming-system.md | 1 + .../streaming-system/wal/tracing.md | 243 +++++++++++ .../replicate_stream_client_impl.go | 49 ++- .../replicate_stream_client_trace_test.go | 130 ++++++ .../streaming/internal/producer/producer.go | 27 +- .../internal/producer/producer_task.go | 40 +- .../internal/producer/producer_trace_test.go | 297 ++++++++++++++ .../pipeline/flow_graph_dd_node.go | 63 +-- .../pipeline/flow_graph_message.go | 12 + .../pipeline/flow_graph_time_tick_node.go | 24 +- .../pipeline/flow_graph_write_node.go | 11 +- .../replicate/replicate_stream_server.go | 7 +- .../replicate_stream_server_trace_test.go | 136 ++++++ internal/querynodev2/pipeline/delete_node.go | 3 +- internal/querynodev2/pipeline/filter_node.go | 3 +- .../querynodev2/pipeline/filter_node_test.go | 50 +++ internal/querynodev2/pipeline/insert_node.go | 21 +- .../broadcaster/ack_callback_scheduler.go | 23 +- .../ack_callback_scheduler_trace_test.go | 82 ++++ .../server/broadcaster/broadcast_manager.go | 1 - .../server/broadcaster/broadcaster_test.go | 2 +- .../server/broadcaster/broadcaster_with_rk.go | 19 +- .../broadcaster_with_rk_trace_test.go | 187 +++++++++ .../broadcaster/pending_broadcast_task.go | 2 + .../pending_broadcast_task_trace_test.go | 61 +++ .../client/handler/consumer/consumer_impl.go | 28 +- .../client/handler/consumer/consumer_test.go | 213 ++++++++++ .../handler/producer/produce_grpc_client.go | 5 +- .../client/handler/producer/producer_impl.go | 2 +- .../client/handler/producer/producer_test.go | 96 +++++ .../flusher/flusherimpl/flusher_components.go | 17 +- .../server/flusher/flusherimpl/wal_flusher.go | 49 +-- .../flusher/flusherimpl/wal_flusher_test.go | 39 ++ .../handler/producer/produce_server.go | 21 +- .../handler/producer/produce_server_test.go | 2 +- .../producer/produce_server_trace_test.go | 104 +++++ .../server/wal/adaptor/scanner_switchable.go | 21 + .../server/wal/adaptor/wal_adaptor.go | 25 +- .../wal/adaptor/wal_adaptor_trace_test.go | 372 +++++++++++++++++ .../server/wal/metricsutil/append.go | 5 +- .../server/wal/metricsutil/wal_write.go | 8 +- .../wal/metricsutil/wal_write_trace_test.go | 92 +++++ .../recovery/recovery_drop_collection_test.go | 21 +- .../recovery_storage_flushall_test.go | 9 +- .../wal/recovery/recovery_storage_impl.go | 138 ++++--- .../wal/recovery/recovery_storage_test.go | 2 +- .../wal/recovery/replicate_checkpoint_test.go | 17 +- .../wal/recovery/salvage_checkpoint_test.go | 10 +- .../util/mock_message/mock_MutableMessage.go | 3 +- pkg/proto/messages.proto | 12 + pkg/proto/messagespb/messages.pb.go | 276 ++++++++----- pkg/streaming/util/message/adaptor/message.go | 25 +- .../util/message/adaptor/message_test.go | 81 ++++ pkg/streaming/util/message/builder.go | 4 + pkg/streaming/util/message/message_impl.go | 13 + pkg/streaming/util/message/properties.go | 1 + pkg/streaming/util/message/trace.go | 201 +++++++++ pkg/streaming/util/message/trace_test.go | 386 ++++++++++++++++++ 60 files changed, 3467 insertions(+), 350 deletions(-) create mode 100644 docs/agent_guides/observability/README.md create mode 100644 docs/agent_guides/streaming-system/wal/tracing.md create mode 100644 internal/cdc/replication/replicatestream/replicate_stream_client_trace_test.go create mode 100644 internal/distributed/streaming/internal/producer/producer_trace_test.go create mode 100644 internal/proxy/replicate/replicate_stream_server_trace_test.go create mode 100644 internal/streamingcoord/server/broadcaster/ack_callback_scheduler_trace_test.go create mode 100644 internal/streamingcoord/server/broadcaster/broadcaster_with_rk_trace_test.go create mode 100644 internal/streamingcoord/server/broadcaster/pending_broadcast_task_trace_test.go create mode 100644 internal/streamingnode/server/service/handler/producer/produce_server_trace_test.go create mode 100644 internal/streamingnode/server/wal/adaptor/wal_adaptor_trace_test.go create mode 100644 internal/streamingnode/server/wal/metricsutil/wal_write_trace_test.go create mode 100644 pkg/streaming/util/message/trace.go create mode 100644 pkg/streaming/util/message/trace_test.go diff --git a/CLAUDE.md b/CLAUDE.md index e56a10cf4e..8536ac7a93 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -32,6 +32,7 @@ NEVER answer based on documentation alone or code alone. NEVER skip Step 2 — t ### Subsystems Reference +- [**Observability**](docs/agent_guides/observability/README.md): Logging, metrics, tracing, and observability debug workflows. - [**Streaming System**](docs/agent_guides/streaming-system/streaming-system.md): Write path, WAL, DDL/DCL execution, replication && CDC. ## Testing diff --git a/docs/agent_guides/observability/README.md b/docs/agent_guides/observability/README.md new file mode 100644 index 0000000000..a658780fb3 --- /dev/null +++ b/docs/agent_guides/observability/README.md @@ -0,0 +1,24 @@ +# Observability - AI Agent Guides + +This directory contains observability guides for AI agents working on Milvus. +Use these guides before changing logging, metrics, tracing, or related +configuration. + +## Guides + +| Guide | Use When | +|---|---| +| [mlog - AI Agent Logging Guide](logging.md) | Adding or changing application logs. Covers `mlog` usage, context requirements, fields, levels, and logging rules. | +| [WAL Tracing](../streaming-system/wal/tracing.md) | Understanding or changing WAL trace span semantics across append, consume, transaction, broadcast, and replication paths. | + +## Rules of Thumb + +- Use `mlog` for all Milvus logs. Do not use `zap`, the old `pkg/log` package, + the standard `log` package, or `fmt.Println` for runtime logging. +- Keep observability hot paths cheap. Avoid payload logging and + high-cardinality metric labels. +- When debugging, start from the narrowest available evidence such as trace ID, + request time window, node, collection, channel, or error message. +- Preserve compatibility for metric names, label sets, config keys, and log + field names unless the task explicitly requires a breaking change. +- Add or update focused tests when changing observability behavior. diff --git a/docs/agent_guides/streaming-system/streaming-system.md b/docs/agent_guides/streaming-system/streaming-system.md index 797e3f1510..dad9afb4a8 100644 --- a/docs/agent_guides/streaming-system/streaming-system.md +++ b/docs/agent_guides/streaming-system/streaming-system.md @@ -33,6 +33,7 @@ Every WAL entry is a [**Message**](message/message.md) representing a system eve - [**Lock**](wal/lock.md): Exclusive/shared append access at VChannel or PChannel scope. - [**Shard Management**](wal/shard-management.md): Per-PChannel collection/partition/segment metadata and segment assignment. - [**RecoveryStorage**](wal/recovery-storage.md): Checkpoint, metadata and data persistence and WAL-based state recovery. +- [**WAL Tracing**](wal/tracing.md): Trace span semantics for append, consume, broadcast, transaction, and replication paths. **[StreamingClient](streaming-client/streaming-client.md)**: In-process Append/Read/Broadcast API with service discovery and auto-reconnect. diff --git a/docs/agent_guides/streaming-system/wal/tracing.md b/docs/agent_guides/streaming-system/wal/tracing.md new file mode 100644 index 0000000000..5371b96f1b --- /dev/null +++ b/docs/agent_guides/streaming-system/wal/tracing.md @@ -0,0 +1,243 @@ +# WAL Tracing + +> **How to use this guide**: This page defines the semantic shape of WAL +> traces. Use it when changing WAL trace spans, message-carried trace context, +> replication tracing, broadcast tracing, or consume-side trace restoration. +> Read the streaming system and message semantic guides first when the change +> also modifies WAL behavior. + +WAL tracing describes the causal path of a logical WAL message through append, +durable persistence, consume, replication, and callback handling. + +It is not a per-function profiler. WAL spans are semantic markers for lifecycle +boundaries. If a span does not represent a WAL ownership, persistence, consume, +replication, or callback boundary, it usually does not belong in the WAL trace. + +## Global Model + +WAL trace context is message-carried. A WAL message stores the trace context +that downstream asynchronous work should use as its parent. When a message +crosses a semantic ownership boundary, the message-carried trace context may be +overwritten to the new parent. + +WAL tracing follows these principles: + +1. A client request may produce one or more WAL messages. +2. A WAL message carries trace context across asynchronous boundaries. +3. Write-side spans describe how a message enters WAL. +4. Append spans describe concrete persistence attempts for concrete messages. +5. Consume-side spans describe where WAL visibility resumes from stored + message state. +6. Replication spans describe secondary-cluster ownership of a primary message. +7. Broadcast callback spans describe follow-up work after broadcast delivery is + acknowledged. +8. TimeTick is intentionally not traced. + +The helper APIs live in the message package. This guide defines their intended +trace semantics, not their implementation. + +## Span Semantics + +| Span | Meaning | Parent | Duration | +|---|---|---|---| +| `wal.autocommit` | Logical WAL write for one non-transactional, non-broadcast message. | Caller request or upstream WAL trace. | Covers the client-side logical append operation. | +| `wal.txn` | Logical WAL write for one transaction, including BeginTxn, body messages, and CommitTxn. | Caller request. | Covers the whole transaction append sequence. | +| `wal.broadcast` | Logical WAL broadcast for one broadcast task across target pchannels or vchannels. | Caller request. | Covers broadcast fan-out and append scheduling. | +| `wal.append` | WAL adaptor append boundary for one concrete message append. | A logical write span such as `wal.autocommit`, `wal.txn`, `wal.broadcast`, `replicate.secondary`, or `wal.dist_append`. | Covers adaptor-level append work. | +| `wal.appendimpl` | WAL implementation append boundary where the concrete backend persists the message. | `wal.append`. | Covers backend append and persistence. | +| `wal.dist_append` | Distributed append marker when a producer writes through a remote WAL. | The logical write span, usually `wal.autocommit` or `wal.txn`. | Covers the remote append request until append completion. | +| `wal.catchup_consume` | Durable backend consume marker for a message read from the local backend scanner during catchup. | Message-carried trace. | Short marker span. | +| `wal.dist_consume` | Distributed or non-local consume marker for a message read through a remote scanner or remote WAL path. | Message-carried trace. | Short marker span. | +| `replicate.secondary` | Secondary cluster receive and re-append boundary for a replicated primary WAL message. | Primary-side consumed message trace, usually under `wal.dist_consume`. | Covers secondary-side replicate handling until append. | +| `wal.bc_callback` | Broadcast ACK callback processing after broadcast message persistence and acknowledgement. | Broadcast message trace. | Covers callback handling such as task completion and cache invalidation. | + +`wal.autocommit`, `wal.txn`, and `wal.broadcast` are logical write roots. They +represent user-visible or system-visible WAL write intent. + +`wal.append` and `wal.appendimpl` are physical append boundaries. They should +not become logical roots unless the upstream context is missing. + +`wal.catchup_consume` and `wal.dist_consume` are resume markers. They are +usually short and exist to reconnect downstream asynchronous work to the +message trace that was stored in WAL. + +`wal.catchup_consume` is emitted only when the scanner reads from the durable +backend scanner in catchup mode. It is intentionally absent in tailing mode: +tailing readers consume the same immutable message instance from the +WriteAheadBuffer, and that shared message's properties must not be mutated to +overwrite `_tc`. + +`replicate.secondary` is the only replication ownership span. There is no +`replicate.primary` span. + +## Span Attributes + +Span attributes should explain the message or broadcast being traced without +encoding that information into span names. Keep span names stable and put +message scope, timing, transaction, and broadcast metadata in attributes. + +Common message attributes: + +| Attribute | Applies To | Meaning | +|---|---|---| +| `message.type` | Message-related WAL spans. | WAL message type. | +| `message.vchannel` | VChannel-scoped messages. | Target VChannel. Empty means the message is PChannel-level or not VChannel-scoped. | +| `message.timetick` | Message-related WAL spans after TimeTick is assigned. | WAL TimeTick of the traced message. This does not make TimeTick messages traceable. | +| `message.replicate` | Message-related WAL spans. | Whether the message carries replication metadata. | +| `txn.id` | Transaction messages and synthetic transaction traces. | Transaction ID. | + +Broadcast-specific attributes on `wal.broadcast`: + +| Attribute | Meaning | +|---|---| +| `broadcast.id` | BroadcastID of the broadcast task. | +| `broadcast.vchannels` | Target broadcast VChannels. | +| `message.type` | Broadcast message type. | + +`wal.broadcast` should make the broadcast target scope visible through +`broadcast.vchannels`, not by changing the span name. + +## Canonical Trace Shapes + +### Autocommit + +Autocommit is the normal path for one non-transactional, non-broadcast message. +A complete trace may include remote append, primary persistence, consume, and +secondary replication: + +```text +request span + wal.autocommit # autocommit-specific + wal.dist_append # remote append only + wal.append + wal.appendimpl + wal.catchup_consume / wal.dist_consume # consume marker, catchup/remote only + replicate.secondary # replication only + wal.autocommit # secondary re-append + wal.append + wal.appendimpl +``` + +If the producer already owns a local WAL, `wal.dist_append` is absent and +`wal.append` is directly under `wal.autocommit`. If the message is consumed from +a local durable backend scanner during catchup, the consume marker is +`wal.catchup_consume`; if it is consumed through a remote or distributed +scanner path, the marker is `wal.dist_consume`. A steady-state local tailing +consumer emits no consume marker. + +The secondary-side `wal.autocommit` is the local append of a replicated concrete +message into the secondary WAL. It does not mean the original client request was +issued on the secondary cluster. + +### Transaction + +A transaction has one transaction-level logical write span and several concrete +message appends. A complete trace uses CommitTxn as the point where the +transaction becomes consumable: + +```text +request span + wal.txn # txn-specific + wal.dist_append # remote append only: BeginTxn + wal.append # BeginTxn + wal.appendimpl + wal.dist_append # remote append only: txn body + wal.append # txn body message + wal.appendimpl + wal.dist_append # remote append only: txn body + wal.append # txn body message + wal.appendimpl + wal.dist_append # remote append only: CommitTxn + wal.append # CommitTxn + wal.appendimpl + wal.dist_consume # consume marker for assembled txn + replicate.secondary # replication only: BeginTxn + wal.autocommit # secondary re-append + wal.append + wal.appendimpl + replicate.secondary # replication only: txn body + wal.autocommit # secondary re-append + wal.append + wal.appendimpl + replicate.secondary # replication only: CommitTxn + wal.autocommit # secondary re-append + wal.append + wal.appendimpl +``` + +If the producer writes to a local WAL, `wal.dist_append` is absent and each +`wal.append` is directly under `wal.txn`. + +`wal.txn` is the semantic parent for the whole transaction. BeginTxn, body +messages, and CommitTxn are each concrete WAL messages, so each append has its +own `wal.append` and `wal.appendimpl`. + +Downstream consumption uses the transaction assembled at CommitTxn. The +synthetic transaction message is the downstream semantic unit. When that +transaction is expanded later, BeginTxn, body messages, and CommitTxn should use +the transaction-level trace rather than preserving unrelated body-level traces. +This means the expanded child messages may have their `_tc` overwritten from +the assembled transaction's CommitTxn trace; repeated copying of that +transaction trace is intentional. + +Txn tracing should stay flat at the logical level. Do not add a separate +`client.append` span or independent per-body logical roots. + +### Broadcast + +Broadcast has one broadcast-level logical root and multiple concrete appends. A +complete trace may include primary broadcast append, primary callback, +distributed consume, secondary re-append, and secondary callback: + +```text +request span + wal.broadcast # broadcast-specific + wal.append + wal.appendimpl + wal.dist_consume # consume marker + replicate.secondary # replication only + wal.append + wal.appendimpl + wal.bc_callback # broadcast-specific callback + wal.append + wal.appendimpl + wal.bc_callback # broadcast-specific callback +``` + +`wal.broadcast` represents the broadcast task. Each `wal.append` represents one +concrete append produced by broadcast fan-out. + +Broadcast must not create `wal.autocommit` or `wal.txn`. Broadcast is already +the logical write root. + +`wal.bc_callback` represents ACK-driven callback work after broadcast delivery, +such as task completion and cache invalidation. It is not a new user request. + +## Non-Traceable Messages + +TimeTick is intentionally not traced. + +TimeTick is a WAL progress and control signal, not a user-visible mutation. +Tracing every TimeTick would dominate trace volume and hide useful message +causality. + +Trace propagation tests must not use TimeTick unless the expected behavior is a +trace no-op. + +## Invariants + +- Span names are stable and low-cardinality. +- WAL tracing is message-causal, not goroutine-causal. +- Message trace context represents the parent for downstream asynchronous work. +- Logical write spans are `wal.autocommit`, `wal.txn`, and `wal.broadcast`. +- Physical append spans are `wal.append` and `wal.appendimpl`. +- Consume spans are short resume markers. +- Replication has only `replicate.secondary`. +- Broadcast does not create `wal.autocommit` or `wal.txn`. +- `wal.broadcast` carries BroadcastID, broadcast VChannels, and message type as + attributes. +- Transaction downstream fan-out uses transaction-level trace. +- Message-related WAL spans carry message type, TimeTick when available, + VChannel when applicable, and replication state. +- Transaction messages carry txn ID when available. +- TimeTick is not traced. diff --git a/internal/cdc/replication/replicatestream/replicate_stream_client_impl.go b/internal/cdc/replication/replicatestream/replicate_stream_client_impl.go index 8bc94741fd..0d2991cb95 100644 --- a/internal/cdc/replication/replicatestream/replicate_stream_client_impl.go +++ b/internal/cdc/replication/replicatestream/replicate_stream_client_impl.go @@ -27,7 +27,6 @@ import ( "google.golang.org/grpc/codes" "google.golang.org/grpc/status" - "github.com/milvus-io/milvus-proto/go-api/v3/commonpb" "github.com/milvus-io/milvus-proto/go-api/v3/milvuspb" "github.com/milvus-io/milvus/internal/cdc/cluster" "github.com/milvus-io/milvus/internal/cdc/meta" @@ -224,25 +223,7 @@ func (r *replicateStreamClient) sendLoop(ctx context.Context) (err error) { } if msg.MessageType() == message.MessageTypeTxn { txnMsg := message.AsImmutableTxnMessage(msg) - - // send txn begin message - beginMsg := txnMsg.Begin() - err := r.sendMessage(beginMsg) - if err != nil { - return err - } - - // send txn messages - err = txnMsg.RangeOver(func(msg message.ImmutableMessage) error { - return r.sendMessage(msg) - }) - if err != nil { - return err - } - - // send txn commit message - commitMsg := txnMsg.Commit() - err = r.sendMessage(commitMsg) + err = r.sendTxnMessage(txnMsg) if err != nil { return err } @@ -257,6 +238,8 @@ func (r *replicateStreamClient) sendLoop(ctx context.Context) (err error) { } func (r *replicateStreamClient) sendMessage(msg message.ImmutableMessage) (err error) { + immutableMessage := msg.IntoImmutableMessageProto() + defer func() { logger := mlog.With(mlog.String("key", r.channel.Key), mlog.Int64("revision", r.channel.ModRevision)) if err != nil { @@ -266,22 +249,36 @@ func (r *replicateStreamClient) sendMessage(msg message.ImmutableMessage) (err e logger.Debug(r.ctx, "send message success", mlog.FieldMessage(msg)) } }() - immutableMessage := msg.IntoImmutableMessageProto() + req := &milvuspb.ReplicateRequest{ Request: &milvuspb.ReplicateRequest_ReplicateMessage{ ReplicateMessage: &milvuspb.ReplicateMessage{ SourceClusterId: r.clusterID, - Message: &commonpb.ImmutableMessage{ - Id: msg.MessageID().IntoProto(), - Payload: immutableMessage.GetPayload(), - Properties: immutableMessage.GetProperties(), - }, + Message: immutableMessage, }, }, } return r.client.Send(req) } +func (r *replicateStreamClient) sendTxnMessage(txnMsg message.ImmutableTxnMessage) (err error) { + // send txn begin message + if err = r.sendMessage(txnMsg.Begin()); err != nil { + return err + } + + // send txn body messages + if err = txnMsg.RangeOver(func(msg message.ImmutableMessage) error { + return r.sendMessage(msg) + }); err != nil { + return err + } + + // send txn commit message + err = r.sendMessage(txnMsg.Commit()) + return +} + func (r *replicateStreamClient) recvLoop(ctx context.Context) (err error) { logger := mlog.With(mlog.String("key", r.channel.Key), mlog.Int64("revision", r.channel.ModRevision)) defer func() { diff --git a/internal/cdc/replication/replicatestream/replicate_stream_client_trace_test.go b/internal/cdc/replication/replicatestream/replicate_stream_client_trace_test.go new file mode 100644 index 0000000000..76d7eda0fd --- /dev/null +++ b/internal/cdc/replication/replicatestream/replicate_stream_client_trace_test.go @@ -0,0 +1,130 @@ +//go:build test && dynamic + +package replicatestream + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "go.opentelemetry.io/otel" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" + "go.opentelemetry.io/otel/trace" + + "github.com/milvus-io/milvus/internal/cdc/meta" + mock_message "github.com/milvus-io/milvus/pkg/v3/mocks/streaming/util/mock_message" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" + "github.com/milvus-io/milvus/pkg/v3/streaming/walimpls/impls/walimplstest" +) + +// buildTestCDCImmutableMessage creates an ImmutableMessage with the trace context from primaryCtx injected. +func buildTestCDCImmutableMessage(t *testing.T, primaryCtx context.Context) message.ImmutableMessage { + t.Helper() + msgID := walimplstest.NewTestMessageID(1) + mutableMsg := message.CreateTestEmptyInsertMesage(1, nil) + mutableMsg.WithTimeTick(100) + mutableMsg.WithLastConfirmed(msgID) + message.InjectTraceContext(primaryCtx, mutableMsg) + return mutableMsg.IntoImmutableMessage(msgID) +} + +func setupTraceExporter(t *testing.T) *tracetest.InMemoryExporter { + t.Helper() + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + t.Cleanup(func() { + otel.SetTracerProvider(prev) + }) + return exporter +} + +// TestSendMessage_PreservesTraceContextWithoutPrimarySpan asserts that CDC +// forwarding does not create a replicate.primary span or rewrite message _tc. +func TestSendMessage_PreservesTraceContextWithoutPrimarySpan(t *testing.T) { + exporter := setupTraceExporter(t) + + // Simulate a primary WAL message with a persisted _tc pointing at + // the primary wal.append span. + primaryCtx, primarySpan := otel.Tracer("test").Start(context.Background(), "primary.wal.append") + primarySC := trace.SpanContextFromContext(primaryCtx) + primarySpan.End() + + imsg := buildTestCDCImmutableMessage(t, primaryCtx) + client := newMockReplicateStreamClient(t) + + c := &replicateStreamClient{ + clusterID: "test-cluster", + client: client, + channel: &meta.ReplicateChannel{Key: "test-replicate-key"}, + metrics: NewReplicateMetrics(nil), + } + err := c.sendMessage(imsg) + assert.NoError(t, err) + capturedReq := <-client.ch + + spans := exporter.GetSpans() + for _, s := range spans { + assert.NotEqual(t, "replicate.primary", s.Name, "CDC send should be represented by wal.catchup_consume") + } + + outMsg := message.MilvusMessageToImmutableMessage(capturedReq.GetReplicateMessage().GetMessage()) + outSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), outMsg)) + assert.True(t, outSC.IsValid(), "outgoing _tc must be valid") + assert.Equal(t, primarySC.TraceID(), outSC.TraceID(), + "outgoing _tc should preserve the immutable message trace ID") + assert.Equal(t, primarySC.SpanID(), outSC.SpanID(), + "outgoing _tc should preserve the immutable message span ID") + assert.Equal(t, primarySC.SpanID(), trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), imsg)).SpanID(), + "sendMessage should not mutate immutable message properties") +} + +// TestSendTxnMessage_SendsEachMessageWithItsOwnCdcSpan verifies that txn +// replication does not add an extra txn-level span. +func TestSendTxnMessage_SendsEachMessageWithItsOwnCdcSpan(t *testing.T) { + exporter := setupTraceExporter(t) + + // Simulate a primary wal.txn trace context in the Begin message. + primaryCtx, primarySpan := otel.Tracer("test").Start(context.Background(), "primary.wal.txn") + primarySC := trace.SpanContextFromContext(primaryCtx) + primarySpan.End() + + beginMsg := buildTestCDCImmutableMessage(t, primaryCtx) + bodyMsg := buildTestCDCImmutableMessage(t, context.Background()) + commitMsg := buildTestCDCImmutableMessage(t, context.Background()) + + // Build a mock ImmutableTxnMessage. + txnMock := mock_message.NewMockImmutableTxnMessage(t) + txnMock.EXPECT().Begin().Return(beginMsg) + txnMock.EXPECT().RangeOver(mock.Anything).RunAndReturn(func(fn func(message.ImmutableMessage) error) error { + return fn(bodyMsg) + }) + txnMock.EXPECT().Commit().Return(commitMsg) + + client := newMockReplicateStreamClient(t) + c := &replicateStreamClient{ + clusterID: "test-cluster", + client: client, + channel: &meta.ReplicateChannel{Key: "test-replicate-key"}, + metrics: NewReplicateMetrics(nil), + } + err := c.sendTxnMessage(txnMock) + assert.NoError(t, err) + + // begin + 1 body + commit = 3 sends. + assert.Len(t, client.ch, 3) + + spans := exporter.GetSpans() + + for _, s := range spans { + assert.NotEqual(t, "replicate.primary.txn", s.Name, "txn replication should not emit a txn-level span") + assert.NotEqual(t, "replicate.primary", s.Name, "txn replication should not emit per-message replicate.primary spans") + } + _ = primarySC +} diff --git a/internal/distributed/streaming/internal/producer/producer.go b/internal/distributed/streaming/internal/producer/producer.go index 286bf265cc..07e7e7081a 100644 --- a/internal/distributed/streaming/internal/producer/producer.go +++ b/internal/distributed/streaming/internal/producer/producer.go @@ -7,10 +7,13 @@ import ( "github.com/cenkalti/backoff/v4" "github.com/cockroachdb/errors" + "go.opentelemetry.io/otel/codes" + "go.opentelemetry.io/otel/trace" "github.com/milvus-io/milvus/internal/distributed/streaming/internal/errs" "github.com/milvus-io/milvus/internal/streamingnode/client/handler" "github.com/milvus-io/milvus/internal/streamingnode/client/handler/producer" + "github.com/milvus-io/milvus/internal/streamingnode/client/handler/registry" "github.com/milvus-io/milvus/internal/util/streamingutil/status" "github.com/milvus-io/milvus/pkg/v3/mlog" "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" @@ -124,7 +127,15 @@ func (p *ResumableProducer) produceInternal(ctx context.Context, msg message.Mut if err != nil { return nil, err } - produceResult, err := producerHandler.Append(ctx, msg) + appendCtx, span := p.startDistAppendSpanIfRemote(ctx, producerHandler, msg) + produceResult, err := producerHandler.Append(appendCtx, msg) + if span != nil { + if err != nil { + span.RecordError(err) + span.SetStatus(codes.Error, err.Error()) + } + span.End() + } if err == nil { return produceResult, nil } @@ -152,6 +163,20 @@ func (p *ResumableProducer) produceInternal(ctx context.Context, msg message.Mut } } +func (p *ResumableProducer) startDistAppendSpanIfRemote(ctx context.Context, producerHandler handler.Producer, msg message.MutableMessage) (context.Context, trace.Span) { + if isLocalProducer(producerHandler) { + return ctx, nil + } + return message.StartSpanForMessage(ctx, msg, message.SpanNameWALDistAppend) +} + +func isLocalProducer(producerHandler handler.Producer) bool { + if pm, ok := producerHandler.(producerWithMetrics); ok { + return registry.IsLocal(pm.Producer) + } + return registry.IsLocal(producerHandler) +} + // resumeLoop is used to resume producer from error. func (p *ResumableProducer) resumeLoop() { defer func() { diff --git a/internal/distributed/streaming/internal/producer/producer_task.go b/internal/distributed/streaming/internal/producer/producer_task.go index f632bd246e..1e5656cb7b 100644 --- a/internal/distributed/streaming/internal/producer/producer_task.go +++ b/internal/distributed/streaming/internal/producer/producer_task.go @@ -21,6 +21,7 @@ import ( "sync" "time" + "go.opentelemetry.io/otel/codes" "golang.org/x/time/rate" "github.com/milvus-io/milvus/internal/util/streamingutil/status" @@ -123,16 +124,47 @@ func (g *ProduceGuard) commit(ctx context.Context) (*types.AppendResult, error) if len(g.msgs) == 0 { panic("append task with no messages") } + if g.msgs[0].BroadcastHeader() != nil { + if len(g.msgs) != 1 { + panic("broadcast guard must hold exactly one message") + } + return g.producer.produceInternal(ctx, g.msgs[0]) + } // auto commit if there's only one message. if len(g.msgs) == 1 { - return g.producer.produceInternal(ctx, g.msgs[0]) + return g.produceAutocommit(ctx, g.msgs[0]) } // produce with transaction. return g.produceTxn(ctx, g.msgs...) } +func (g *ProduceGuard) produceAutocommit(ctx context.Context, msg message.MutableMessage) (_ *types.AppendResult, err error) { + ctx, span := message.StartSpanForMessage(ctx, msg, message.SpanNameWALAutocommit) + defer func() { + if err != nil { + span.RecordError(err) + span.SetStatus(codes.Error, err.Error()) + } + span.End() + }() + message.InjectTraceContext(ctx, msg) + return g.producer.produceInternal(ctx, msg) +} + // produceTxn produces the messages with a transaction, retry if the transaction is expired. -func (g *ProduceGuard) produceTxn(ctx context.Context, msgs ...message.MutableMessage) (*types.AppendResult, error) { +func (g *ProduceGuard) produceTxn(ctx context.Context, msgs ...message.MutableMessage) (_ *types.AppendResult, err error) { + ctx, span := message.StartSpan(ctx, message.SpanNameWALTxn) + defer func() { + if err != nil { + span.RecordError(err) + span.SetStatus(codes.Error, err.Error()) + } + span.End() + }() + for _, msg := range msgs { + message.InjectTraceContext(ctx, msg) + } + for { if ctx.Err() != nil { return nil, ctx.Err() @@ -153,7 +185,7 @@ func (g *ProduceGuard) produceTxn(ctx context.Context, msgs ...message.MutableMe } // produceWithTxnOnce produces the messages with a transaction once. -func (g *ProduceGuard) produceWithTxnOnce(ctx context.Context, msgs ...message.MutableMessage) (*types.AppendResult, error) { +func (g *ProduceGuard) produceWithTxnOnce(ctx context.Context, msgs ...message.MutableMessage) (_ *types.AppendResult, err error) { // a txn batch should always belong to one vchannel. txn, err := g.beginTxn(ctx, msgs[0].VChannel()) if err != nil { @@ -173,6 +205,7 @@ func (g *ProduceGuard) beginTxn(ctx context.Context, vchannel string) (*message. WithHeader(&message.BeginTxnMessageHeader{}). WithBody(&message.BeginTxnMessageBody{}). MustBuildMutable() + message.InjectTraceContext(ctx, beginTxn) result, err := g.producer.produceInternal(ctx, beginTxn) if err != nil { @@ -215,6 +248,7 @@ func (g *ProduceGuard) commitTxn(ctx context.Context, vchannel string, txn *mess WithHeader(&message.CommitTxnMessageHeader{}). WithBody(&message.CommitTxnMessageBody{}). MustBuildMutable() + message.InjectTraceContext(ctx, commitTxn) return g.producer.produceInternal(ctx, commitTxn.WithTxnContext(*txn)) } diff --git a/internal/distributed/streaming/internal/producer/producer_trace_test.go b/internal/distributed/streaming/internal/producer/producer_trace_test.go new file mode 100644 index 0000000000..b9d3c94dfc --- /dev/null +++ b/internal/distributed/streaming/internal/producer/producer_trace_test.go @@ -0,0 +1,297 @@ +//go:build test && dynamic + +package producer + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/bytedance/mockey" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "go.opentelemetry.io/otel" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" + "go.opentelemetry.io/otel/trace" + + "github.com/milvus-io/milvus-proto/go-api/v3/msgpb" + "github.com/milvus-io/milvus/internal/mocks/streamingnode/client/handler/mock_producer" + "github.com/milvus-io/milvus/internal/streamingnode/client/handler" + "github.com/milvus-io/milvus/internal/streamingnode/client/handler/producer" + "github.com/milvus-io/milvus/pkg/v3/mocks/streaming/util/mock_message" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/types" +) + +func TestProduceAutocommit_OpensSpanAndInjectsTraceContext(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + p := newTestResumableProducer(t) + msg := buildTestInsertMessage(t) + + g := &ProduceGuard{ + producer: p, + msgs: []message.MutableMessage{msg}, + } + _, err := g.commit(context.Background()) + assert.NoError(t, err) + + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), msg)) + spans := exporter.GetSpans() + var autocommit tracetest.SpanStub + for _, s := range spans { + assert.NotEqual(t, "wal.append.client", s.Name, "client append span should be flattened") + if s.Name == message.SpanNameWALAutocommit { + autocommit = s + } + } + assert.Equal(t, message.SpanNameWALAutocommit, autocommit.Name, "wal.autocommit span should be emitted") + assert.True(t, msgSC.IsValid(), "autocommit message should carry _tc") + assert.Equal(t, autocommit.SpanContext.TraceID(), msgSC.TraceID()) + assert.Equal(t, autocommit.SpanContext.SpanID(), msgSC.SpanID()) +} + +// buildTestInsertMessage builds a minimal MutableMessage for testing. +func buildTestInsertMessage(t *testing.T) message.MutableMessage { + t.Helper() + return message.NewInsertMessageBuilderV1(). + WithHeader(&message.InsertMessageHeader{ + CollectionId: 1, + }). + WithBody(&msgpb.InsertRequest{ + CollectionID: 1, + }). + WithVChannel("test-vchannel"). + MustBuildMutable() +} + +func TestProduceTxn_WrapsInWalTxnSpan(t *testing.T) { + defer mockey.UnPatchAll() + + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + // Stub the three inner steps of produceWithTxnOnce so we only exercise + // the span-wrapping logic of the function itself. + mockey.Mock((*ProduceGuard).beginTxn).Return(&message.TxnContext{}, nil).Build() + mockey.Mock((*ProduceGuard).appendTxnBody).Return(nil).Build() + mockey.Mock((*ProduceGuard).commitTxn).Return(&types.AppendResult{}, nil).Build() + + g := &ProduceGuard{} + msg := buildTestInsertMessage(t) + + _, err := g.produceTxn(context.Background(), msg) + assert.NoError(t, err) + + spans := exporter.GetSpans() + var walTxn tracetest.SpanStub + for _, s := range spans { + assert.NotEqual(t, "wal.append.client", s.Name, "client append span should be flattened") + if s.Name == message.SpanNameWALTxn { + walTxn = s + break + } + } + assert.Equal(t, message.SpanNameWALTxn, walTxn.Name, "wal.txn span should be emitted") +} + +func TestProduceTxn_UsesSameTraceContextForTxnMessages(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + mockMsgID := mock_message.NewMockMessageID(t) + txnCtx := &message.TxnContext{ + TxnID: 1, + Keepalive: time.Minute, + } + + var mu sync.Mutex + var msgSpanContexts []trace.SpanContext + mockProd := mock_producer.NewMockProducer(t) + mockProd.EXPECT().Append(mock.Anything, mock.Anything). + RunAndReturn(func(_ context.Context, msg message.MutableMessage) (*types.AppendResult, error) { + sc := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), msg)) + mu.Lock() + msgSpanContexts = append(msgSpanContexts, sc) + mu.Unlock() + return &types.AppendResult{ + MessageID: mockMsgID, + TimeTick: 100, + TxnCtx: txnCtx, + }, nil + }) + mockProd.EXPECT().Available().Return(make(chan struct{})) + mockProd.EXPECT().IsAvailable().Return(true) + mockProd.EXPECT().Close().Return() + + rp := NewResumableProducer(func(ctx context.Context, opts *handler.ProducerOptions) (producer.Producer, error) { + return mockProd, nil + }, &ProducerOptions{ + PChannel: "test-trace", + }) + t.Cleanup(rp.Close) + + g := &ProduceGuard{producer: rp} + _, err := g.produceTxn(context.Background(), buildTestInsertMessage(t), buildTestInsertMessage(t)) + assert.NoError(t, err) + + spans := exporter.GetSpans() + var walTxn tracetest.SpanStub + for _, s := range spans { + assert.NotEqual(t, "wal.append.client", s.Name, "client append span should be flattened") + if s.Name == message.SpanNameWALTxn { + walTxn = s + break + } + } + assert.Equal(t, message.SpanNameWALTxn, walTxn.Name, "wal.txn span should be emitted") + assert.Len(t, msgSpanContexts, 4, "begin, two body messages and commit should be appended") + for _, sc := range msgSpanContexts { + assert.True(t, sc.IsValid(), "txn message should carry _tc") + assert.Equal(t, walTxn.SpanContext.TraceID(), sc.TraceID()) + assert.Equal(t, walTxn.SpanContext.SpanID(), sc.SpanID()) + } +} + +func TestProduceBroadcast_DoesNotOpenAutocommitOrTxnSpan(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + p := newTestResumableProducer(t) + msg := message.NewDropCollectionMessageBuilderV1(). + WithHeader(&message.DropCollectionMessageHeader{}). + WithBody(&msgpb.DropCollectionRequest{}). + WithBroadcast([]string{"test-vchannel"}). + MustBuildBroadcast(). + WithBroadcastID(1). + SplitIntoMutableMessage()[0] + + g := &ProduceGuard{ + producer: p, + msgs: []message.MutableMessage{msg}, + } + _, err := g.commit(context.Background()) + assert.NoError(t, err) + + for _, s := range exporter.GetSpans() { + assert.NotEqual(t, message.SpanNameWALAutocommit, s.Name, "broadcast append should not emit wal.autocommit") + assert.NotEqual(t, message.SpanNameWALTxn, s.Name, "broadcast append should not emit wal.txn") + assert.NotEqual(t, "wal.append.client", s.Name, "client append span should be flattened") + } +} + +func TestProduceBroadcast_PanicsWhenGuardHasMultipleMessages(t *testing.T) { + msgs := message.NewDropCollectionMessageBuilderV1(). + WithHeader(&message.DropCollectionMessageHeader{}). + WithBody(&msgpb.DropCollectionRequest{}). + WithBroadcast([]string{"test-vchannel-1", "test-vchannel-2"}). + MustBuildBroadcast(). + WithBroadcastID(1). + SplitIntoMutableMessage() + + g := &ProduceGuard{ + msgs: msgs, + } + assert.PanicsWithValue(t, "broadcast guard must hold exactly one message", func() { + _, _ = g.commit(context.Background()) + }) +} + +func TestProduceInternalRemoteOpensDistAppendSpan(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + mockMsgID := mock_message.NewMockMessageID(t) + msg := buildTestInsertMessage(t) + + mockProd := mock_producer.NewMockProducer(t) + mockProd.EXPECT().Append(mock.Anything, mock.Anything). + RunAndReturn(func(ctx context.Context, msg message.MutableMessage) (*types.AppendResult, error) { + sc := trace.SpanContextFromContext(ctx) + assert.True(t, sc.IsValid(), "remote append ctx should carry wal.dist_append span") + return &types.AppendResult{ + MessageID: mockMsgID, + TimeTick: 100, + }, nil + }) + mockProd.EXPECT().Available().Return(make(chan struct{})) + mockProd.EXPECT().IsAvailable().Return(true) + mockProd.EXPECT().Close().Return() + + rp := NewResumableProducer(func(ctx context.Context, opts *handler.ProducerOptions) (producer.Producer, error) { + return mockProd, nil + }, &ProducerOptions{ + PChannel: "test-trace", + }) + t.Cleanup(rp.Close) + + _, err := rp.produceInternal(context.Background(), msg) + assert.NoError(t, err) + + spans := exporter.GetSpans() + var distAppend tracetest.SpanStub + for _, s := range spans { + if s.Name == message.SpanNameWALDistAppend { + distAppend = s + break + } + } + assert.Equal(t, message.SpanNameWALDistAppend, distAppend.Name, "wal.dist_append span should be emitted for remote append") +} + +// newTestResumableProducer builds a ResumableProducer with a mocked inner handler +// so that produceInternal exits on the first iteration. +func newTestResumableProducer(t *testing.T) *ResumableProducer { + t.Helper() + + mockMsgID := mock_message.NewMockMessageID(t) + mockProd := mock_producer.NewMockProducer(t) + mockProd.EXPECT().Append(mock.Anything, mock.Anything).Return(&types.AppendResult{ + MessageID: mockMsgID, + TimeTick: 100, + }, nil) + mockProd.EXPECT().Available().Return(make(chan struct{})) + mockProd.EXPECT().IsAvailable().Return(true) + mockProd.EXPECT().Close().Return() + + rp := NewResumableProducer(func(ctx context.Context, opts *handler.ProducerOptions) (producer.Producer, error) { + return mockProd, nil + }, &ProducerOptions{ + PChannel: "test-trace", + }) + t.Cleanup(rp.Close) + return rp +} diff --git a/internal/flushcommon/pipeline/flow_graph_dd_node.go b/internal/flushcommon/pipeline/flow_graph_dd_node.go index 69287a32c2..4946f00a33 100644 --- a/internal/flushcommon/pipeline/flow_graph_dd_node.go +++ b/internal/flushcommon/pipeline/flow_graph_dd_node.go @@ -145,10 +145,11 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { } for _, msg := range msMsg.TsMessages() { + msgCtx := msg.TraceCtx() switch msg.Type() { case commonpb.MsgType_DropCollection: if msg.(*msgstream.DropCollectionMsg).GetCollectionID() == ddn.collectionID { - mlog.Info(ddn.ctx, "Receiving DropCollection msg", mlog.String("channel", ddn.vChannelName)) + mlog.Info(msgCtx, "Receiving DropCollection msg", mlog.String("channel", ddn.vChannelName)) ddn.dropMode.Store(true) fgMsg.dropCollection = true } @@ -156,14 +157,14 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { case commonpb.MsgType_DropPartition: dpMsg := msg.(*msgstream.DropPartitionMsg) if dpMsg.GetCollectionID() == ddn.collectionID { - mlog.Info(ddn.ctx, "drop partition msg received", mlog.String("channel", ddn.vChannelName), mlog.FieldPartitionID(dpMsg.GetPartitionID())) + mlog.Info(msgCtx, "drop partition msg received", mlog.String("channel", ddn.vChannelName), mlog.FieldPartitionID(dpMsg.GetPartitionID())) fgMsg.dropPartitions = append(fgMsg.dropPartitions, dpMsg.PartitionID) } case commonpb.MsgType_Insert: imsg := msg.(*msgstream.InsertMsg) if imsg.CollectionID != ddn.collectionID { - mlog.Warn(ddn.ctx, "filter invalid insert message, collection mis-match", + mlog.Warn(msgCtx, "filter invalid insert message, collection mis-match", mlog.Int64("Get collID", imsg.CollectionID), mlog.String("channel", ddn.vChannelName), mlog.Int64("Expected collID", ddn.collectionID)) @@ -171,7 +172,7 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { } if ddn.tryToFilterSegmentInsertMessages(imsg) { - mlog.Debug(ddn.ctx, "filter insert messages", + mlog.Debug(msgCtx, "filter insert messages", mlog.Int64("filter segmentID", imsg.GetSegmentID()), mlog.String("channel", ddn.vChannelName), mlog.Uint64("message timestamp", msg.EndTs()), @@ -193,7 +194,7 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { WithLabelValues(paramtable.GetStringNodeID(), metrics.InsertLabel). Add(float64(imsg.GetNumRows())) - mlog.Debug(ddn.ctx, "DDNode receive insert messages", + mlog.Debug(msgCtx, "DDNode receive insert messages", mlog.FieldSegmentID(imsg.GetSegmentID()), mlog.String("channel", ddn.vChannelName), mlog.Int("numRows", len(imsg.GetRowIDs())), @@ -205,14 +206,14 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { dmsg := msg.(*msgstream.DeleteMsg) if dmsg.CollectionID != ddn.collectionID { - mlog.Warn(ddn.ctx, "filter invalid DeleteMsg, collection mis-match", + mlog.Warn(msgCtx, "filter invalid DeleteMsg, collection mis-match", mlog.Int64("Get collID", dmsg.CollectionID), mlog.String("channel", ddn.vChannelName), mlog.Int64("Expected collID", ddn.collectionID)) continue } - mlog.Debug(ddn.ctx, "DDNode receive delete messages", + mlog.Debug(msgCtx, "DDNode receive delete messages", mlog.String("channel", ddn.vChannelName), mlog.Int64("numRows", dmsg.NumRows), mlog.Uint64("startPosTs", msMsg.StartPositions()[0].GetTimestamp()), @@ -239,11 +240,11 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { mlog.Int32("msgType", int32(msg.Type())), mlog.Uint64("timetick", createSegment.CreateSegmentMessage.TimeTick()), ) - logger.Info(ddn.ctx, "receive create segment message") - if err := ddn.msgHandler.HandleCreateSegment(ddn.ctx, createSegment.CreateSegmentMessage); err != nil { - logger.Warn(ddn.ctx, "handle create segment message failed", mlog.Err(err)) + logger.Info(msgCtx, "receive create segment message") + if err := ddn.msgHandler.HandleCreateSegment(msgCtx, createSegment.CreateSegmentMessage); err != nil { + logger.Warn(msgCtx, "handle create segment message failed", mlog.Err(err)) } else { - logger.Info(ddn.ctx, "handle create segment message success") + logger.Info(msgCtx, "handle create segment message success") } case commonpb.MsgType_FlushSegment: flushMsg := msg.(*adaptor.FlushMessageBody) @@ -252,11 +253,11 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { mlog.Int32("msgType", int32(msg.Type())), mlog.Uint64("timetick", flushMsg.FlushMessage.TimeTick()), ) - logger.Info(ddn.ctx, "receive flush message") + logger.Info(msgCtx, "receive flush message") if err := ddn.msgHandler.HandleFlush(flushMsg.FlushMessage); err != nil { - logger.Warn(ddn.ctx, "handle flush message failed", mlog.Err(err)) + logger.Warn(msgCtx, "handle flush message failed", mlog.Err(err)) } else { - logger.Info(ddn.ctx, "handle flush message success") + logger.Info(msgCtx, "handle flush message success") } case commonpb.MsgType_ManualFlush: manualFlushMsg := msg.(*adaptor.ManualFlushMessageBody) @@ -267,15 +268,15 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { mlog.Uint64("flushTs", manualFlushMsg.ManualFlushMessage.Header().FlushTs), mlog.Int64s("segmentIDs", manualFlushMsg.ManualFlushMessage.Header().SegmentIds), ) - logger.Info(ddn.ctx, "receive manual flush message") + logger.Info(msgCtx, "receive manual flush message") if err := ddn.msgHandler.HandleManualFlush(manualFlushMsg.ManualFlushMessage); err != nil { - logger.Warn(ddn.ctx, "handle manual flush message failed", mlog.Err(err)) + logger.Warn(msgCtx, "handle manual flush message failed", mlog.Err(err)) } else { - logger.Info(ddn.ctx, "handle manual flush message success") + logger.Info(msgCtx, "handle manual flush message success") } case commonpb.MsgType_FlushAll: flushAllMsg := msg.(*adaptor.FlushAllMessageBody) - mlog.Info(ddn.ctx, "receive flush all message", + mlog.Info(msgCtx, "receive flush all message", mlog.FieldVChannel(ddn.Name()), mlog.Int32("msgType", int32(msg.Type())), mlog.Uint64("timetick", flushAllMsg.FlushAllMessage.TimeTick()), @@ -289,8 +290,8 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { mlog.Uint64("timetick", schemaMsg.SchemaChangeMessage.TimeTick()), mlog.Int64s("segmentIDs", schemaMsg.SchemaChangeMessage.Header().FlushedSegmentIds), ) - logger.Info(ddn.ctx, "receive schema change message") - ddn.msgHandler.HandleSchemaChange(ddn.ctx, schemaMsg.SchemaChangeMessage) + logger.Info(msgCtx, "receive schema change message") + ddn.msgHandler.HandleSchemaChange(msgCtx, schemaMsg.SchemaChangeMessage) case commonpb.MsgType_AlterCollection: alterCollectionMsg := msg.(*adaptor.AlterCollectionMessageBody) logger := mlog.With( @@ -298,11 +299,11 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { mlog.Int32("msgType", int32(msg.Type())), mlog.Uint64("timetick", alterCollectionMsg.AlterCollectionMessage.TimeTick()), ) - logger.Info(ddn.ctx, "receive put collection message") - if err := ddn.msgHandler.HandleAlterCollection(ddn.ctx, alterCollectionMsg.AlterCollectionMessage); err != nil { - logger.Warn(ddn.ctx, "handle put collection message failed", mlog.Err(err)) + logger.Info(msgCtx, "receive put collection message") + if err := ddn.msgHandler.HandleAlterCollection(msgCtx, alterCollectionMsg.AlterCollectionMessage); err != nil { + logger.Warn(msgCtx, "handle put collection message failed", mlog.Err(err)) } else { - logger.Info(ddn.ctx, "handle put collection message success") + logger.Info(msgCtx, "handle put collection message success") } case commonpb.MsgType_TruncateCollection: truncateCollectionMsg := msg.(*adaptor.TruncateCollectionMessageBody) @@ -312,11 +313,11 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { mlog.Uint64("timetick", truncateCollectionMsg.TruncateCollectionMessage.TimeTick()), mlog.Int64s("segmentIDs", truncateCollectionMsg.TruncateCollectionMessage.Header().SegmentIds), ) - logger.Info(ddn.ctx, "receive truncate collection message") + logger.Info(msgCtx, "receive truncate collection message") if err := ddn.msgHandler.HandleTruncateCollection(truncateCollectionMsg.TruncateCollectionMessage); err != nil { - logger.Warn(ddn.ctx, "handle truncate collection message failed", mlog.Err(err)) + logger.Warn(msgCtx, "handle truncate collection message failed", mlog.Err(err)) } else { - logger.Info(ddn.ctx, "handle truncate collection message success") + logger.Info(msgCtx, "handle truncate collection message success") } case commonpb.MsgType_AlterWAL: alterWALMsg := msg.(*adaptor.AlterWALMessageBody) @@ -326,11 +327,11 @@ func (ddn *ddNode) Operate(in []Msg) []Msg { mlog.Stringer("targetWalName", alterWALMsg.AlterWALMessage.Header().TargetWalName), mlog.Uint64("timetick", alterWALMsg.AlterWALMessage.TimeTick()), ) - logger.Info(ddn.ctx, "receive alter wal message") - if err := ddn.msgHandler.HandleAlterWAL(ddn.ctx, alterWALMsg.AlterWALMessage, ddn.vChannelName); err != nil { - logger.Warn(ddn.ctx, "handle alter wal message failed", mlog.Err(err)) + logger.Info(msgCtx, "receive alter wal message") + if err := ddn.msgHandler.HandleAlterWAL(msgCtx, alterWALMsg.AlterWALMessage, ddn.vChannelName); err != nil { + logger.Warn(msgCtx, "handle alter wal message failed", mlog.Err(err)) } else { - logger.Info(ddn.ctx, "handle alter wal message success") + logger.Info(msgCtx, "handle alter wal message success") } fgMsg.isAlterWal = true fgMsg.alterWalTimeTick = alterWALMsg.AlterWALMessage.TimeTick() diff --git a/internal/flushcommon/pipeline/flow_graph_message.go b/internal/flushcommon/pipeline/flow_graph_message.go index 4ab912fd13..5e086d7f2b 100644 --- a/internal/flushcommon/pipeline/flow_graph_message.go +++ b/internal/flushcommon/pipeline/flow_graph_message.go @@ -17,6 +17,8 @@ package pipeline import ( + "context" + "github.com/milvus-io/milvus-proto/go-api/v3/msgpb" "github.com/milvus-io/milvus/internal/flushcommon/util" "github.com/milvus-io/milvus/internal/flushcommon/writebuffer" @@ -70,6 +72,16 @@ func (fgMsg *FlowGraphMsg) TimeTick() typeutil.Timestamp { return fgMsg.TimeRange.TimestampMax } +func (fgMsg *FlowGraphMsg) TraceCtx() context.Context { + if len(fgMsg.InsertMessages) > 0 { + return fgMsg.InsertMessages[0].TraceCtx() + } + if len(fgMsg.DeleteMessages) > 0 { + return fgMsg.DeleteMessages[0].TraceCtx() + } + return context.Background() +} + func (fgMsg *FlowGraphMsg) IsClose() bool { return fgMsg.IsCloseMsg() } diff --git a/internal/flushcommon/pipeline/flow_graph_time_tick_node.go b/internal/flushcommon/pipeline/flow_graph_time_tick_node.go index cbd553983c..3ad70b1756 100644 --- a/internal/flushcommon/pipeline/flow_graph_time_tick_node.go +++ b/internal/flushcommon/pipeline/flow_graph_time_tick_node.go @@ -74,6 +74,7 @@ func (ttn *ttNode) Close() { // Operate handles input messages, implementing flowgraph.Node func (ttn *ttNode) Operate(in []Msg) []Msg { fgMsg := in[0].(*FlowGraphMsg) + ctx := fgMsg.TraceCtx() if fgMsg.dropCollection { ttn.dropMode.Store(true) if ttn.dropCallback != nil { @@ -88,7 +89,7 @@ func (ttn *ttNode) Operate(in []Msg) []Msg { // skip updating checkpoint for drop collection // even if its the close msg if ttn.dropMode.Load() { - mlog.RatedInfo(context.TODO(), rate.Limit(1.0), "ttnode in dropMode", mlog.String("channel", ttn.vChannelName)) + mlog.RatedInfo(ctx, rate.Limit(1.0), "ttnode in dropMode", mlog.String("channel", ttn.vChannelName)) return []Msg{} } @@ -101,13 +102,13 @@ func (ttn *ttNode) Operate(in []Msg) []Msg { if len(fgMsg.EndPositions) > 0 { channelPos, _, err := ttn.writeBufferManager.GetCheckpoint(ttn.vChannelName) if err != nil { - mlog.Warn(context.TODO(), "channel removed", mlog.String("channel", ttn.vChannelName), mlog.Err(err)) + mlog.Warn(ctx, "channel removed", mlog.String("channel", ttn.vChannelName), mlog.Err(err)) return []Msg{} } - mlog.Info(context.TODO(), "flowgraph is closing, force update channel CP", + mlog.Info(ctx, "flowgraph is closing, force update channel CP", mlog.Time("cpTs", tsoutil.PhysicalTime(channelPos.GetTimestamp())), mlog.String("channel", channelPos.GetChannelName())) - ttn.updateChannelCP(channelPos, curTs, false) + ttn.updateChannelCP(ctx, channelPos, curTs, false) } return in } @@ -116,27 +117,26 @@ func (ttn *ttNode) Operate(in []Msg) []Msg { channelPos, needUpdate, err := ttn.writeBufferManager.GetCheckpoint(ttn.vChannelName) if fgMsg.isAlterWal && !needUpdate { - channelPos, needUpdate, err = ttn.waitForCheckpointUpdate(fgMsg, curTs) + channelPos, needUpdate, err = ttn.waitForCheckpointUpdate(ctx, fgMsg, curTs) } if err != nil { - mlog.Warn(context.TODO(), "channel removed", mlog.String("channel", ttn.vChannelName), mlog.Err(err)) + mlog.Warn(ctx, "channel removed", mlog.String("channel", ttn.vChannelName), mlog.Err(err)) return []Msg{} } if curTs.Sub(ttn.lastUpdateTime.Load()) >= paramtable.Get().DataNodeCfg.UpdateChannelCheckpointInterval.GetAsDuration(time.Second) { - ttn.updateChannelCP(channelPos, curTs, false) + ttn.updateChannelCP(ctx, channelPos, curTs, false) return []Msg{} } if needUpdate { - ttn.updateChannelCP(channelPos, curTs, true) + ttn.updateChannelCP(ctx, channelPos, curTs, true) } return []Msg{} } // waitForCheckpointUpdate waits for checkpoint to be ready using exponential backoff retry -func (ttn *ttNode) waitForCheckpointUpdate(fgMsg *FlowGraphMsg, curTs time.Time) (*msgpb.MsgPosition, bool, error) { - ctx := context.Background() +func (ttn *ttNode) waitForCheckpointUpdate(ctx context.Context, fgMsg *FlowGraphMsg, curTs time.Time) (*msgpb.MsgPosition, bool, error) { backoffConfig := backoff.NewExponentialBackOff() backoffConfig.InitialInterval = 100 * time.Millisecond backoffConfig.MaxInterval = 1 * time.Second @@ -191,12 +191,12 @@ func (ttn *ttNode) waitForCheckpointUpdate(fgMsg *FlowGraphMsg, curTs time.Time) return channelPos, needUpdate, nil } -func (ttn *ttNode) updateChannelCP(channelPos *msgpb.MsgPosition, curTs time.Time, flush bool) { +func (ttn *ttNode) updateChannelCP(ctx context.Context, channelPos *msgpb.MsgPosition, curTs time.Time, flush bool) { callBack := func() { channelCPTs, _ := tsoutil.ParseTS(channelPos.GetTimestamp()) // reset flush ts to prevent frequent flush ttn.writeBufferManager.NotifyCheckpointUpdated(ttn.vChannelName, channelPos.GetTimestamp()) - mlog.Debug(context.TODO(), "UpdateChannelCheckpoint success", + mlog.Debug(ctx, "UpdateChannelCheckpoint success", mlog.String("channel", ttn.vChannelName), mlog.Uint64("cpTs", channelPos.GetTimestamp()), mlog.Stringer("walName", channelPos.WALName), diff --git a/internal/flushcommon/pipeline/flow_graph_write_node.go b/internal/flushcommon/pipeline/flow_graph_write_node.go index ee0911a483..9f7e720f1e 100644 --- a/internal/flushcommon/pipeline/flow_graph_write_node.go +++ b/internal/flushcommon/pipeline/flow_graph_write_node.go @@ -91,11 +91,12 @@ func (wNode *writeNode) Operate(in []Msg) []Msg { }() start, end := fgMsg.StartPositions[0], fgMsg.EndPositions[0] + ctx := fgMsg.TraceCtx() currentSchema := wNode.metacache.GetSchema(fgMsg.TimeTick()) schemaVersion := currentSchema.GetVersion() functionOutputFieldIDs, err := wNode.functionStore.OutputFieldIDs(currentSchema) if err != nil { - mlog.Error(context.TODO(), "failed to get embedding output fields", mlog.Err(err)) + mlog.Error(ctx, "failed to get embedding output fields", mlog.Err(err)) panic(err) } @@ -106,13 +107,13 @@ func (wNode *writeNode) Operate(in []Msg) []Msg { continue } if err := wNode.functionStore.FillEmbeddingData(wNode.collectionID, currentSchema, msg.InsertRequest); err != nil { - mlog.Error(context.TODO(), "failed to fill embedding data", mlog.Err(err)) + mlog.Error(msg.TraceCtx(), "failed to fill embedding data", mlog.Err(err)) panic(err) } } preparedInsertData, err := writebuffer.PrepareInsert(currentSchema, wNode.pkField, fgMsg.InsertMessages) if err != nil { - mlog.Error(context.TODO(), "failed to prepare data", mlog.Err(err)) + mlog.Error(ctx, "failed to prepare data", mlog.Err(err)) panic(err) } insertData = preparedInsertData @@ -121,7 +122,7 @@ func (wNode *writeNode) Operate(in []Msg) []Msg { err = wNode.wbManager.BufferData(wNode.channelName, fgMsg.InsertData, fgMsg.DeleteMessages, start, end, schemaVersion) if err != nil { - mlog.Error(context.TODO(), "failed to buffer data", mlog.Err(err)) + mlog.Error(ctx, "failed to buffer data", mlog.Err(err)) panic(err) } @@ -130,7 +131,7 @@ func (wNode *writeNode) Operate(in []Msg) []Msg { func(id int64, _ int) (*commonpb.SegmentStats, bool) { segInfo, ok := wNode.metacache.GetSegmentByID(id) if !ok { - mlog.Warn(context.TODO(), "segment not found for stats", mlog.Int64("segment", id)) + mlog.Warn(ctx, "segment not found for stats", mlog.Int64("segment", id)) return nil, false } return &commonpb.SegmentStats{ diff --git a/internal/proxy/replicate/replicate_stream_server.go b/internal/proxy/replicate/replicate_stream_server.go index 141f660c72..c9c1c2a568 100644 --- a/internal/proxy/replicate/replicate_stream_server.go +++ b/internal/proxy/replicate/replicate_stream_server.go @@ -119,8 +119,12 @@ func (p *ReplicateStreamServer) handleReplicateMessage(req *milvuspb.ReplicateRe if err != nil { return err } + ctx := message.ExtractTraceContext(p.streamServer.Context(), msg) + ctx, span := message.StartSpanForMessage(ctx, msg, message.SpanNameReplicateSecondary) + message.OverwriteTraceContext(ctx, msg) + defer span.End() + sourceTs := msg.ReplicateHeader().TimeTick - ctx := p.streamServer.Context() mlog.Debug(ctx, "recv replicate message from client", mlog.String("messageID", reqMsg.GetId().GetId()), mlog.Uint64("sourceTimeTick", sourceTs), @@ -138,6 +142,7 @@ func (p *ReplicateStreamServer) handleReplicateMessage(req *milvuspb.ReplicateRe p.sendReplicateResult(sourceTs, msg) return nil } + span.RecordError(err) // unexpected error, will close the stream and wait for client to reconnect. mlog.Warn(ctx, "append replicate message to wal failed", mlog.FieldMessage(msg), mlog.Err(err)) return err diff --git a/internal/proxy/replicate/replicate_stream_server_trace_test.go b/internal/proxy/replicate/replicate_stream_server_trace_test.go new file mode 100644 index 0000000000..67f369785b --- /dev/null +++ b/internal/proxy/replicate/replicate_stream_server_trace_test.go @@ -0,0 +1,136 @@ +//go:build test && dynamic + +package replicate + +import ( + "context" + "testing" + + "github.com/apache/pulsar-client-go/pulsar" + "github.com/bytedance/mockey" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "go.opentelemetry.io/otel" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" + "go.opentelemetry.io/otel/trace" + + "github.com/milvus-io/milvus-proto/go-api/v3/commonpb" + "github.com/milvus-io/milvus-proto/go-api/v3/milvuspb" + "github.com/milvus-io/milvus-proto/go-api/v3/msgpb" + "github.com/milvus-io/milvus/internal/distributed/streaming" + "github.com/milvus-io/milvus/internal/mocks/distributed/mock_streaming" + "github.com/milvus-io/milvus/pkg/v3/proto/messagespb" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/types" + pulsar2 "github.com/milvus-io/milvus/pkg/v3/streaming/walimpls/impls/pulsar" +) + +// TestHandleReplicateMessage_OpensWalReplicateAppendSpan verifies that +// handleReplicateMessage extracts the replicated message trace context, +// opens a "replicate.secondary" child span, and overwrites the local mutable +// message trace context before append. +func TestHandleReplicateMessage_OpensWalReplicateAppendSpan(t *testing.T) { + defer mockey.UnPatchAll() + + // Set up an in-memory OTel exporter and make it the global provider. + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + // Build a source-side traced context and capture the expected trace ID. + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), "source.wal.append") + sourceSC := trace.SpanContextFromContext(sourceCtx) + expectedTraceID := sourceSC.TraceID() + sourceSpan.End() + + // Build a replicate message proto with _tc carried by the immutable message. + reqMsg := buildTraceTestReplicateMsgProto(t, sourceCtx) + + req := &milvuspb.ReplicateRequest_ReplicateMessage{ + ReplicateMessage: &milvuspb.ReplicateMessage{ + SourceClusterId: "cluster-a", + Message: reqMsg, + }, + } + + // Capture the ctx passed to Append. + var capturedCtx context.Context + var capturedMsg message.ReplicateMutableMessage + + replicateService := mock_streaming.NewMockReplicateService(t) + replicateService.EXPECT().Append(mock.Anything, mock.Anything). + RunAndReturn(func(ctx context.Context, msg message.ReplicateMutableMessage) (*types.AppendResult, error) { + capturedCtx = ctx + capturedMsg = msg + return &types.AppendResult{TimeTick: 1}, nil + }) + mockWAL := mock_streaming.NewMockWALAccesser(t) + mockWAL.EXPECT().Replicate().Return(replicateService) + streaming.SetWALForTest(mockWAL) + + // Build a minimal ReplicateStreamServer using the existing package helper. + ctx := createContextWithClusterID("cluster-a") + mockStream := newMockReplicateStreamServer(ctx) + server, err := CreateReplicateServer(mockStream) + assert.NoError(t, err) + + // Call handleReplicateMessage directly (synchronous). + err = server.handleReplicateMessage(req) + assert.NoError(t, err) + + // Flush the provider to ensure all spans are exported. + _ = tp.ForceFlush(context.Background()) + + // Assert that a "replicate.secondary" span was emitted with the right trace ID. + spans := exporter.GetSpans() + var walSpan tracetest.SpanStub + for _, s := range spans { + if s.Name == message.SpanNameReplicateSecondary { + walSpan = s + assert.Equal(t, expectedTraceID, s.SpanContext.TraceID(), + "replicate.secondary span must share the source trace ID") + } + } + assert.Equal(t, message.SpanNameReplicateSecondary, walSpan.Name, "a 'replicate.secondary' span must be emitted") + assert.Equal(t, sourceSC.SpanID(), walSpan.Parent.SpanID(), + "replicate.secondary should be a child of the source message span") + + // Also verify that the ctx passed to Append carries the same trace ID. + if capturedCtx != nil { + capturedSpan := trace.SpanFromContext(capturedCtx) + assert.True(t, capturedSpan.SpanContext().IsValid(), + "ctx passed to Append should carry a valid span") + assert.Equal(t, expectedTraceID, capturedSpan.SpanContext().TraceID(), + "ctx passed to Append must share the source trace ID") + } + + assert.NotNil(t, capturedMsg) + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), capturedMsg)) + assert.True(t, msgSC.IsValid(), "replicate server should overwrite the mutable message trace context") + assert.Equal(t, walSpan.SpanContext.TraceID(), msgSC.TraceID()) + assert.Equal(t, walSpan.SpanContext.SpanID(), msgSC.SpanID()) +} + +// buildTraceTestReplicateMsgProto builds a *commonpb.ImmutableMessage that +// carries _tc through the normal message conversion path. +func buildTraceTestReplicateMsgProto(t *testing.T, tracedCtx context.Context) *commonpb.ImmutableMessage { + t.Helper() + messageID := pulsar2.NewPulsarID(pulsar.EarliestMessageID()) + tt := uint64(42) + msg := message.NewInsertMessageBuilderV1(). + WithVChannel("test-vchannel"). + WithHeader(&messagespb.InsertMessageHeader{}). + WithBody(&msgpb.InsertRequest{}). + MustBuildMutable().WithTimeTick(tt). + WithLastConfirmed(messageID) + + message.InjectTraceContext(tracedCtx, msg) + milvusMsg := message.ImmutableMessageToMilvusMessage(commonpb.WALName_Pulsar.String(), msg.IntoImmutableMessage(messageID)) + return milvusMsg +} diff --git a/internal/querynodev2/pipeline/delete_node.go b/internal/querynodev2/pipeline/delete_node.go index 0bef09eaa9..bdfa041d99 100644 --- a/internal/querynodev2/pipeline/delete_node.go +++ b/internal/querynodev2/pipeline/delete_node.go @@ -41,6 +41,7 @@ type deleteNode struct { // addDeleteData find the segment of delete column in DeleteMsg and save in deleteData func (dNode *deleteNode) addDeleteData(deleteDatas map[UniqueID]*delegator.DeleteData, msg *DeleteMsg) { + ctx := msg.TraceCtx() deleteData, ok := deleteDatas[msg.PartitionID] if !ok { deleteData = &delegator.DeleteData{ @@ -53,7 +54,7 @@ func (dNode *deleteNode) addDeleteData(deleteDatas map[UniqueID]*delegator.Delet deleteData.Timestamps = append(deleteData.Timestamps, msg.Timestamps...) deleteData.RowCount += int64(len(pks)) - mlog.Info(context.TODO(), "pipeline fetch delete msg", + mlog.Info(ctx, "pipeline fetch delete msg", mlog.FieldCollectionID(dNode.collectionID), mlog.FieldPartitionID(msg.PartitionID), mlog.Int("deleteRowNum", len(pks)), diff --git a/internal/querynodev2/pipeline/filter_node.go b/internal/querynodev2/pipeline/filter_node.go index 0aed0df819..f6a97a29ed 100644 --- a/internal/querynodev2/pipeline/filter_node.go +++ b/internal/querynodev2/pipeline/filter_node.go @@ -85,9 +85,10 @@ func (fNode *filterNode) Operate(in Msg) Msg { // add msg to out if msg pass check of filter for _, msg := range streamMsgPack.Msgs { + ctx := msg.TraceCtx() err := fNode.filtrate(collection, msg) if err != nil { - mlog.Debug(context.TODO(), "filter invalid message", + mlog.Debug(ctx, "filter invalid message", mlog.String("message type", msg.Type().String()), mlog.String("channel", fNode.channel), mlog.FieldCollectionID(fNode.collectionID), diff --git a/internal/querynodev2/pipeline/filter_node_test.go b/internal/querynodev2/pipeline/filter_node_test.go index 17aa3cd7e8..2e771f642f 100644 --- a/internal/querynodev2/pipeline/filter_node_test.go +++ b/internal/querynodev2/pipeline/filter_node_test.go @@ -17,11 +17,14 @@ package pipeline import ( + "context" "testing" "github.com/samber/lo" "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" + "go.opentelemetry.io/otel/trace" "github.com/milvus-io/milvus/internal/querynodev2/delegator" "github.com/milvus-io/milvus/internal/querynodev2/segments" @@ -199,3 +202,50 @@ func (suite *FilterNodeSuite) buildMsgPack() *msgstream.MsgPack { func TestFilterNode(t *testing.T) { suite.Run(t, new(FilterNodeSuite)) } + +func TestFilterNodePreservesTraceContext(t *testing.T) { + paramtable.Init() + + expectedTraceID, err := trace.TraceIDFromHex("0102030405060708090a0b0c0d0e0f10") + require.NoError(t, err) + expectedSpanID, err := trace.SpanIDFromHex("0102030405060708") + require.NoError(t, err) + clientCtx := trace.ContextWithSpanContext(context.Background(), trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: expectedTraceID, + SpanID: expectedSpanID, + })) + + const ( + collectionID = int64(111) + partitionID = int64(11) + segmentID = int64(3) + channel = "test-channel" + ) + + collection := segments.NewTestCollection(collectionID, querypb.LoadType_LoadCollection, nil) + collection.AddPartition(partitionID) + mockCollectionManager := segments.NewMockCollectionManager(t) + mockCollectionManager.EXPECT().Get(collectionID).Return(collection) + + mockDelegator := delegator.NewMockShardDelegator(t) + mockDelegator.EXPECT().VerifyExcludedSegments(segmentID, mock.Anything).Return(true) + mockDelegator.EXPECT().TryCleanExcludedSegments(mock.Anything) + + node := newFilterNode(collectionID, channel, &segments.Manager{ + Collection: mockCollectionManager, + Segment: segments.NewMockSegmentManager(t), + }, mockDelegator, 8) + + insertMsg := buildInsertMsg(collectionID, partitionID, segmentID, channel, 1) + insertMsg.SetTraceCtx(clientCtx) + out := node.Operate(&msgstream.MsgPack{ + BeginTs: 1, + EndTs: 1, + Msgs: []msgstream.TsMsg{insertMsg}, + }) + require.NotNil(t, out) + + sc := trace.SpanContextFromContext(insertMsg.TraceCtx()) + require.Equal(t, expectedTraceID, sc.TraceID()) + require.Equal(t, expectedSpanID, sc.SpanID()) +} diff --git a/internal/querynodev2/pipeline/insert_node.go b/internal/querynodev2/pipeline/insert_node.go index 3fd313b48f..74215ac33e 100644 --- a/internal/querynodev2/pipeline/insert_node.go +++ b/internal/querynodev2/pipeline/insert_node.go @@ -17,7 +17,6 @@ package pipeline import ( - "context" "fmt" "sort" @@ -46,11 +45,12 @@ type insertNode struct { } func (iNode *insertNode) addInsertData(insertDatas map[UniqueID]*delegator.InsertData, msg *InsertMsg, collection *Collection) { + ctx := msg.TraceCtx() schema := collection.Schema() insertRecord, skippedFields, err := storage.TransferInsertMsgToInsertRecord(schema, msg) if err != nil { err = merr.Wrap(err, "failed to get primary keys") - mlog.Error(context.TODO(), err.Error(), mlog.Int64("collectionID", iNode.collectionID), mlog.String("channel", iNode.channel)) + mlog.Error(ctx, err.Error(), mlog.Int64("collectionID", iNode.collectionID), mlog.String("channel", iNode.channel)) panic(err) } if len(skippedFields) > 0 { @@ -60,7 +60,7 @@ func (iNode *insertNode) addInsertData(insertDatas map[UniqueID]*delegator.Inser // micro-batcher refuses to merge non-Insert/Delete messages. If that // pack-granularity invariant is ever relaxed, filtering against the current // schema could silently drop fields that exist in a newer schema. - mlog.Warn(context.TODO(), "skip insert payload fields absent from current schema, fields are dropped since the message was written", + mlog.Warn(ctx, "skip insert payload fields absent from current schema, fields are dropped since the message was written", mlog.FieldCollectionID(iNode.collectionID), mlog.FieldSegmentID(msg.SegmentID), mlog.String("channel", iNode.channel), @@ -84,27 +84,27 @@ func (iNode *insertNode) addInsertData(insertDatas map[UniqueID]*delegator.Inser } else { err := typeutil.MergeFieldData(iData.InsertRecord.FieldsData, insertRecord.FieldsData) if err != nil { - mlog.Error(context.TODO(), "failed to merge field data", mlog.String("channel", iNode.channel), mlog.Err(err)) + mlog.Error(ctx, "failed to merge field data", mlog.String("channel", iNode.channel), mlog.Err(err)) panic(err) } iData.InsertRecord.NumRows += insertRecord.NumRows } if err := iNode.appendBM25Stats(iData, msg, schema); err != nil { - mlog.Error(context.TODO(), "failed to append BM25 stats from insert message", mlog.String("channel", iNode.channel), mlog.Err(err)) + mlog.Error(ctx, "failed to append BM25 stats from insert message", mlog.String("channel", iNode.channel), mlog.Err(err)) panic(err) } pks, err := segments.GetPrimaryKeys(msg, schema) if err != nil { - mlog.Error(context.TODO(), "failed to get primary keys from insert message", mlog.Err(err)) + mlog.Error(ctx, "failed to get primary keys from insert message", mlog.Err(err)) panic(err) } iData.PrimaryKeys = append(iData.PrimaryKeys, pks...) iData.RowIDs = append(iData.RowIDs, msg.RowIDs...) iData.Timestamps = append(iData.Timestamps, msg.Timestamps...) - mlog.Debug(context.TODO(), "pipeline fetch insert msg", + mlog.Debug(ctx, "pipeline fetch insert msg", mlog.Int64("collectionID", iNode.collectionID), mlog.Int64("segmentID", msg.SegmentID), mlog.Int("insertRowNum", len(pks)), @@ -125,16 +125,17 @@ func (iNode *insertNode) Operate(in Msg) Msg { sort.Slice(nodeMsg.insertMsgs, func(i, j int) bool { return nodeMsg.insertMsgs[i].BeginTs() < nodeMsg.insertMsgs[j].BeginTs() }) + ctx := nodeMsg.insertMsgs[0].TraceCtx() collection := iNode.manager.Collection.Get(iNode.collectionID) if collection == nil { - mlog.Error(context.TODO(), "insertNode with collection not exist", mlog.Int64("collection", iNode.collectionID)) + mlog.Error(ctx, "insertNode with collection not exist", mlog.Int64("collection", iNode.collectionID)) panic("insertNode with collection not exist") } schema := collection.Schema() functionOutputFieldIDs, err := iNode.functionStore.OutputFieldIDs(schema) if err != nil { - mlog.Error(context.TODO(), "failed to get embedding output fields", mlog.String("channel", iNode.channel), mlog.Err(err)) + mlog.Error(ctx, "failed to get embedding output fields", mlog.String("channel", iNode.channel), mlog.Err(err)) panic(err) } @@ -142,7 +143,7 @@ func (iNode *insertNode) Operate(in Msg) Msg { for _, msg := range nodeMsg.insertMsgs { if len(functionOutputFieldIDs) > 0 && !function.HasAllFieldDataByID(msg.GetFieldsData(), functionOutputFieldIDs) { if err := iNode.functionStore.FillEmbeddingData(iNode.collectionID, schema, msg.InsertRequest); err != nil { - mlog.Error(context.TODO(), "failed to fill embedding data for insert message", mlog.String("channel", iNode.channel), mlog.Err(err)) + mlog.Error(msg.TraceCtx(), "failed to fill embedding data for insert message", mlog.String("channel", iNode.channel), mlog.Err(err)) panic(err) } } diff --git a/internal/streamingcoord/server/broadcaster/ack_callback_scheduler.go b/internal/streamingcoord/server/broadcaster/ack_callback_scheduler.go index 98b564e383..b65650df2b 100644 --- a/internal/streamingcoord/server/broadcaster/ack_callback_scheduler.go +++ b/internal/streamingcoord/server/broadcaster/ack_callback_scheduler.go @@ -8,6 +8,7 @@ import ( "time" "github.com/cenkalti/backoff/v4" + "go.opentelemetry.io/otel/codes" "github.com/milvus-io/milvus/internal/streamingcoord/server/broadcaster/registry" "github.com/milvus-io/milvus/pkg/v3/mlog" @@ -273,9 +274,11 @@ func (s *ackCallbackScheduler) doAckCallback(bt *broadcastTask, g *lockGuards) ( TimeTick: result.TimeTick, } } - // call the ack callback until done. + // call the ack callback until done, under the persisted trace context. bt.ObserveAckCallbackBegin() - if err := s.callMessageAckCallbackUntilDone(s.notifier.Context(), msg, makeMap); err != nil { + if err := runAckCallbackWithTrace(s.notifier.Context(), msg, func(spanCtx context.Context) error { + return s.callMessageAckCallbackUntilDone(spanCtx, msg, makeMap) + }); err != nil { return err } bt.ObserveAckCallbackDone() @@ -316,6 +319,22 @@ func (s *ackCallbackScheduler) callMessageAckCallbackUntilDone(ctx context.Conte } } +// runAckCallbackWithTrace extracts the persisted trace context from the +// broadcast task's message Properties and opens a wal.bc_callback +// span under it, invoking fn with the new ctx. Span is always ended, +// and errors are recorded on the span. +func runAckCallbackWithTrace(baseCtx context.Context, msg message.BroadcastMutableMessage, fn func(ctx context.Context) error) error { + parentCtx := message.ExtractTraceContext(baseCtx, msg) + ctx, span := message.StartSpanForMessage(parentCtx, msg, message.SpanNameWALBCCallback) + defer span.End() + err := fn(ctx) + if err != nil { + span.RecordError(err) + span.SetStatus(codes.Error, err.Error()) + } + return err +} + // sortByControlChannelTimeTick sorts the tasks by the time tick of the control channel. func sortByControlChannelTimeTick(tasks []*broadcastTask) { sort.Slice(tasks, func(i, j int) bool { diff --git a/internal/streamingcoord/server/broadcaster/ack_callback_scheduler_trace_test.go b/internal/streamingcoord/server/broadcaster/ack_callback_scheduler_trace_test.go new file mode 100644 index 0000000000..a825571cf7 --- /dev/null +++ b/internal/streamingcoord/server/broadcaster/ack_callback_scheduler_trace_test.go @@ -0,0 +1,82 @@ +//go:build test + +package broadcaster + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "go.opentelemetry.io/otel" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" + "go.opentelemetry.io/otel/trace" + + "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" +) + +func TestRunAckCallbackWithTrace_OpensChildSpan(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + // Simulate a broadcast message with a persisted _tc pointing at the + // original DDL caller's trace. + originCtx, originSpan := otel.Tracer("test").Start(context.Background(), "ddl.caller") + originTraceID := trace.SpanContextFromContext(originCtx).TraceID() + originSpan.End() + + msg := buildTestBroadcastMessageForTrace(t) + message.InjectTraceContext(originCtx, msg) + + var innerCalled bool + var childTraceID trace.TraceID + err := runAckCallbackWithTrace(context.Background(), msg, func(spanCtx context.Context) error { + innerCalled = true + childTraceID = trace.SpanContextFromContext(spanCtx).TraceID() + return nil + }) + + assert.NoError(t, err) + assert.True(t, innerCalled) + assert.Equal(t, originTraceID, childTraceID, "ack callback span must share trace id with persisted parent") + + spans := exporter.GetSpans() + var found bool + for _, s := range spans { + if s.Name == message.SpanNameWALBCCallback { + assert.Equal(t, originTraceID, s.SpanContext.TraceID()) + assertSpanAttribute(t, s.Attributes, "message.type", message.MessageTypeDropCollection.String()) + assertSpanInt64Attribute(t, s.Attributes, "broadcast.id", 0) + assertSpanStringSliceAttribute(t, s.Attributes, "broadcast.vchannels", []string{"v1", "v2"}) + found = true + } + } + assert.True(t, found, "wal.bc_callback span must be emitted") +} + +func TestRunAckCallbackWithTrace_PreservesBaseCancellation(t *testing.T) { + originCtx, originSpan := otel.Tracer("test").Start(context.Background(), "ddl.caller") + originTraceID := trace.SpanContextFromContext(originCtx).TraceID() + originSpan.End() + + msg := buildTestBroadcastMessageForTrace(t) + message.InjectTraceContext(originCtx, msg) + + baseCtx, cancel := context.WithCancel(context.Background()) + cancel() + + var childTraceID trace.TraceID + err := runAckCallbackWithTrace(baseCtx, msg, func(spanCtx context.Context) error { + childTraceID = trace.SpanContextFromContext(spanCtx).TraceID() + return spanCtx.Err() + }) + + assert.ErrorIs(t, err, context.Canceled) + assert.Equal(t, originTraceID, childTraceID, "ack callback span must still use the persisted trace") +} diff --git a/internal/streamingcoord/server/broadcaster/broadcast_manager.go b/internal/streamingcoord/server/broadcaster/broadcast_manager.go index 6ad662e0f6..5ae7a6c386 100644 --- a/internal/streamingcoord/server/broadcaster/broadcast_manager.go +++ b/internal/streamingcoord/server/broadcaster/broadcast_manager.go @@ -293,7 +293,6 @@ func (bm *broadcastTaskManager) Close() { // addBroadcastTask adds the broadcast task into the manager. func (bm *broadcastTaskManager) addBroadcastTask(msg message.BroadcastMutableMessage, broadcastID uint64, guards *lockGuards) *broadcastTask { - msg = msg.OverwriteBroadcastHeader(broadcastID, guards.ResourceKeys()...) newIncomingTask := newBroadcastTaskFromBroadcastMessage(msg, bm.metrics, bm.ackScheduler) newIncomingTask.SetLogger(bm.Logger()) newIncomingTask.WithResourceKeyLockGuards(guards) diff --git a/internal/streamingcoord/server/broadcaster/broadcaster_test.go b/internal/streamingcoord/server/broadcaster/broadcaster_test.go index c52429a95a..75969a82e7 100644 --- a/internal/streamingcoord/server/broadcaster/broadcaster_test.go +++ b/internal/streamingcoord/server/broadcaster/broadcaster_test.go @@ -153,7 +153,7 @@ func TestBroadcaster(t *testing.T) { bc.Close() broadcastAPI, err = bc.WithResourceKeys(context.Background()) assert.NoError(t, err) - _, err = broadcastAPI.Broadcast(context.Background(), nil) + _, err = broadcastAPI.Broadcast(context.Background(), createNewBroadcastMsg([]string{"v1"})) assert.Error(t, err) err = bc.Ack(context.Background(), mock_message.NewMockImmutableMessage(t)) assert.Error(t, err) diff --git a/internal/streamingcoord/server/broadcaster/broadcaster_with_rk.go b/internal/streamingcoord/server/broadcaster/broadcaster_with_rk.go index 97d3a7c202..b20c82fc84 100644 --- a/internal/streamingcoord/server/broadcaster/broadcaster_with_rk.go +++ b/internal/streamingcoord/server/broadcaster/broadcaster_with_rk.go @@ -3,6 +3,8 @@ package broadcaster import ( "context" + "go.opentelemetry.io/otel/codes" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" "github.com/milvus-io/milvus/pkg/v3/streaming/util/types" ) @@ -14,10 +16,23 @@ type broadcasterWithRK struct { } func (b *broadcasterWithRK) Broadcast(ctx context.Context, msg message.BroadcastMutableMessage) (*types.BroadcastAppendResult, error) { - // consume the guards after the broadcast is called to avoid double unlock. + // Consume the guards before handing them to broadcast to avoid double unlock. guards := b.guards b.guards = nil - return b.broadcaster.broadcast(ctx, msg, b.broadcastID, guards) + msg = msg.OverwriteBroadcastHeader(b.broadcastID, guards.ResourceKeys()...) + ctx, span := message.StartSpanForMessage(ctx, msg, message.SpanNameWALBroadcast) + defer span.End() + + // Keep a trace context in the broadcast message so that the DDL ack callback + // can still extract it after the original caller span is long gone. + message.InjectTraceContext(ctx, msg) + + result, err := b.broadcaster.broadcast(ctx, msg, b.broadcastID, guards) + if err != nil { + span.RecordError(err) + span.SetStatus(codes.Error, err.Error()) + } + return result, err } func (b *broadcasterWithRK) Close() { diff --git a/internal/streamingcoord/server/broadcaster/broadcaster_with_rk_trace_test.go b/internal/streamingcoord/server/broadcaster/broadcaster_with_rk_trace_test.go new file mode 100644 index 0000000000..ea0b6de797 --- /dev/null +++ b/internal/streamingcoord/server/broadcaster/broadcaster_with_rk_trace_test.go @@ -0,0 +1,187 @@ +//go:build test && dynamic + +package broadcaster + +import ( + "context" + "testing" + + "github.com/bytedance/mockey" + "github.com/stretchr/testify/assert" + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" + "go.opentelemetry.io/otel/trace" + + "github.com/milvus-io/milvus-proto/go-api/v3/msgpb" + "github.com/milvus-io/milvus/pkg/v3/proto/messagespb" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/types" +) + +func TestBroadcasterWithRK_InjectsTraceContextBeforeTaskPersist(t *testing.T) { + defer mockey.UnPatchAll() + + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + // Stub the inner broadcast call to capture the msg Properties after injection. + var capturedMsg message.BroadcastMutableMessage + resourceKey := message.NewExclusiveCollectionNameResourceKey("db", "collection") + mockey.Mock((*broadcastTaskManager).broadcast).To( + func(_ *broadcastTaskManager, _ context.Context, msg message.BroadcastMutableMessage, _ uint64, _ *lockGuards) (*types.BroadcastAppendResult, error) { + capturedMsg = msg + return &types.BroadcastAppendResult{}, nil + }).Build() + + msg := buildTestBroadcastMessageForTrace(t) + + // Caller ctx carries a traceable span. + ctx, span := otel.Tracer("test").Start(context.Background(), "caller.ddl") + expectedTraceID := trace.SpanContextFromContext(ctx).TraceID() + defer span.End() + + b := &broadcasterWithRK{ + broadcaster: &broadcastTaskManager{}, + broadcastID: 11, + guards: buildTestLockGuards(resourceKey), + } + _, err := b.Broadcast(ctx, msg) + assert.NoError(t, err) + + spans := exporter.GetSpans() + var broadcastSpan tracetest.SpanStub + for _, s := range spans { + if s.Name == message.SpanNameWALBroadcast { + broadcastSpan = s + break + } + } + assert.Equal(t, message.SpanNameWALBroadcast, broadcastSpan.Name, "wal.broadcast span should be emitted") + assert.Equal(t, expectedTraceID, broadcastSpan.SpanContext.TraceID()) + assertSpanAttribute(t, broadcastSpan.Attributes, "message.type", message.MessageTypeDropCollection.String()) + assertSpanInt64Attribute(t, broadcastSpan.Attributes, "broadcast.id", 11) + assertSpanStringSliceAttribute(t, broadcastSpan.Attributes, "broadcast.vchannels", []string{"v1", "v2"}) + + // Verify _tc was injected on the msg observed by the inner broadcast call. + sc := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), capturedMsg)) + assert.True(t, sc.IsValid(), "_tc should be present after Broadcast") + assert.Equal(t, broadcastSpan.SpanContext.TraceID(), sc.TraceID()) + assert.Equal(t, broadcastSpan.SpanContext.SpanID(), sc.SpanID()) + assert.Equal(t, uint64(11), capturedMsg.BroadcastHeader().BroadcastID) + assert.True(t, capturedMsg.BroadcastHeader().ResourceKeys.Contain(resourceKey)) +} + +func TestBroadcasterWithRK_KeepsExistingTraceContext(t *testing.T) { + defer mockey.UnPatchAll() + + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + var capturedMsg message.BroadcastMutableMessage + mockey.Mock((*broadcastTaskManager).broadcast).To( + func(_ *broadcastTaskManager, _ context.Context, msg message.BroadcastMutableMessage, _ uint64, _ *lockGuards) (*types.BroadcastAppendResult, error) { + capturedMsg = msg + return &types.BroadcastAppendResult{}, nil + }).Build() + + msg := buildTestBroadcastMessageForTrace(t) + originCtx, originSpan := otel.Tracer("test").Start(context.Background(), "origin.ddl") + originSC := trace.SpanContextFromContext(originCtx) + originSpan.End() + message.InjectTraceContext(originCtx, msg) + + callerCtx, callerSpan := otel.Tracer("test").Start(context.Background(), "caller.ddl") + defer callerSpan.End() + + b := &broadcasterWithRK{ + broadcaster: &broadcastTaskManager{}, + broadcastID: 11, + guards: buildTestLockGuards(message.NewExclusiveCollectionNameResourceKey("db", "collection")), + } + _, err := b.Broadcast(callerCtx, msg) + assert.NoError(t, err) + + spans := exporter.GetSpans() + var broadcastSpan tracetest.SpanStub + for _, s := range spans { + if s.Name == message.SpanNameWALBroadcast { + broadcastSpan = s + break + } + } + assert.Equal(t, message.SpanNameWALBroadcast, broadcastSpan.Name, "wal.broadcast span should be emitted") + + sc := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), capturedMsg)) + assert.True(t, sc.IsValid(), "_tc should still be present after Broadcast") + assert.Equal(t, originSC.TraceID(), sc.TraceID()) + assert.Equal(t, originSC.SpanID(), sc.SpanID()) +} + +// buildTestBroadcastMessageForTrace builds a minimal BroadcastMutableMessage for tests. +func buildTestBroadcastMessageForTrace(t *testing.T) message.BroadcastMutableMessage { + t.Helper() + msg, err := message.NewDropCollectionMessageBuilderV1(). + WithHeader(&messagespb.DropCollectionMessageHeader{}). + WithBody(&msgpb.DropCollectionRequest{}). + WithBroadcast([]string{"v1", "v2"}). + BuildBroadcast() + if err != nil { + t.Fatalf("failed to build broadcast message: %v", err) + } + return msg.OverwriteBroadcastHeader(0) +} + +func buildTestLockGuards(keys ...message.ResourceKey) *lockGuards { + guards := &lockGuards{} + for _, key := range keys { + guards.append(&lockGuard{key: key}) + } + return guards +} + +func assertSpanAttribute(t *testing.T, attrs []attribute.KeyValue, key string, value string) { + t.Helper() + for _, attr := range attrs { + if string(attr.Key) == key { + assert.Equal(t, value, attr.Value.AsString()) + return + } + } + t.Fatalf("missing span attribute %q", key) +} + +func assertSpanInt64Attribute(t *testing.T, attrs []attribute.KeyValue, key string, value int64) { + t.Helper() + for _, attr := range attrs { + if string(attr.Key) == key { + assert.Equal(t, value, attr.Value.AsInt64()) + return + } + } + t.Fatalf("missing span attribute %q", key) +} + +func assertSpanStringSliceAttribute(t *testing.T, attrs []attribute.KeyValue, key string, value []string) { + t.Helper() + for _, attr := range attrs { + if string(attr.Key) == key { + assert.ElementsMatch(t, value, attr.Value.AsStringSlice()) + return + } + } + t.Fatalf("missing span attribute %q", key) +} diff --git a/internal/streamingcoord/server/broadcaster/pending_broadcast_task.go b/internal/streamingcoord/server/broadcaster/pending_broadcast_task.go index f50d89e87b..e820b4be46 100644 --- a/internal/streamingcoord/server/broadcaster/pending_broadcast_task.go +++ b/internal/streamingcoord/server/broadcaster/pending_broadcast_task.go @@ -48,6 +48,8 @@ type pendingBroadcastTask struct { // Execute can be repeated called until the task is done. // Same semantics as the `Poll` operation in eventloop. func (b *pendingBroadcastTask) Execute(ctx context.Context) error { + ctx = message.ExtractTraceContext(ctx, b.msg) + if err := b.InitializeRecovery(ctx); err != nil { b.Logger().Warn(ctx, "broadcast task initialize recovery failed", mlog.Err(err)) return err diff --git a/internal/streamingcoord/server/broadcaster/pending_broadcast_task_trace_test.go b/internal/streamingcoord/server/broadcaster/pending_broadcast_task_trace_test.go new file mode 100644 index 0000000000..31148daaba --- /dev/null +++ b/internal/streamingcoord/server/broadcaster/pending_broadcast_task_trace_test.go @@ -0,0 +1,61 @@ +//go:build test + +package broadcaster + +import ( + "context" + "testing" + + "github.com/cockroachdb/errors" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/trace" + + "github.com/milvus-io/milvus/internal/distributed/streaming" + "github.com/milvus-io/milvus/pkg/v3/mlog" + "github.com/milvus-io/milvus/pkg/v3/proto/streamingpb" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/types" +) + +type traceCaptureWALAccesser struct { + streaming.WALAccesser + traceID trace.TraceID +} + +func (w *traceCaptureWALAccesser) AppendMessages(ctx context.Context, msgs ...message.MutableMessage) types.AppendResponses { + w.traceID = trace.SpanContextFromContext(ctx).TraceID() + resps := types.NewAppendResponseN(len(msgs)) + resps.FillAllError(errors.New("append failed")) + return resps +} + +func TestPendingBroadcastTaskExecuteRestoresTraceContextFromMessage(t *testing.T) { + originTraceID, err := trace.TraceIDFromHex("0102030405060708090a0b0c0d0e0f10") + require.NoError(t, err) + originSpanID, err := trace.SpanIDFromHex("0102030405060708") + require.NoError(t, err) + originCtx := trace.ContextWithSpanContext(context.Background(), trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: originTraceID, + SpanID: originSpanID, + })) + + msg := createNewBroadcastMsg([]string{"v1", "v2"}).WithBroadcastID(100) + message.InjectTraceContext(originCtx, msg) + proto := createNewWaitAckBroadcastTaskFromMessage( + msg, + streamingpb.BroadcastTaskState_BROADCAST_TASK_STATE_PENDING, + []byte{0x00, 0x00}, + ) + task := newBroadcastTaskFromProto(proto, newBroadcasterMetrics(), newAckCallbackScheduler(mlog.With())) + task.SetLogger(mlog.With()) + + wal := &traceCaptureWALAccesser{} + oldWAL := streaming.WAL() + streaming.SetWALForTest(wal) + defer streaming.SetWALForTest(oldWAL) + + err = newPendingBroadcastTask(task).Execute(context.Background()) + assert.ErrorIs(t, err, errBroadcastTaskIsNotDone) + assert.Equal(t, originTraceID, wal.traceID) +} diff --git a/internal/streamingnode/client/handler/consumer/consumer_impl.go b/internal/streamingnode/client/handler/consumer/consumer_impl.go index fe619b161c..36073f9e96 100644 --- a/internal/streamingnode/client/handler/consumer/consumer_impl.go +++ b/internal/streamingnode/client/handler/consumer/consumer_impl.go @@ -172,16 +172,18 @@ func (c *consumerImpl) recvLoop() (err error) { resp.Consume.GetMessage().GetPayload(), resp.Consume.GetMessage().GetProperties(), ) + msgCtx := message.ExtractTraceContext(c.ctx, newImmutableMsg) if newImmutableMsg.TxnContext() != nil { - if err := c.handleTxnMessage(newImmutableMsg); err != nil { + if err := c.handleTxnMessage(msgCtx, newImmutableMsg); err != nil { return err } } else { if c.txnBuilder != nil { panic("unreachable code: txn builder should be nil if we receive a non-txn message") } + msgCtx = startDistConsumeSpanForMessage(msgCtx, newImmutableMsg) if result := c.msgHandler.Handle(message.HandleParam{ - Ctx: c.ctx, + Ctx: msgCtx, Message: newImmutableMsg, }); result.Error != nil { c.logger.Warn(c.ctx, "message handle canceled", mlog.Err(err)) @@ -222,7 +224,7 @@ func (c *consumerImpl) createVChannelConsumer() error { return nil } -func (c *consumerImpl) handleTxnMessage(msg message.ImmutableMessage) error { +func (c *consumerImpl) handleTxnMessage(ctx context.Context, msg message.ImmutableMessage) error { switch msg.MessageType() { case message.MessageTypeBeginTxn: if c.txnBuilder != nil { @@ -250,8 +252,10 @@ func (c *consumerImpl) handleTxnMessage(msg message.ImmutableMessage) error { c.logger.Warn(c.ctx, "failed to build txn message", mlog.Any("messageID", commitMsg.MessageID()), mlog.Err(err)) return nil } + ctx = startDistConsumeSpanForMessage(ctx, msg) + overwriteTxnMessagesTraceContext(ctx, message.AsImmutableTxnMessage(msg)) if result := c.msgHandler.Handle(message.HandleParam{ - Ctx: c.ctx, + Ctx: ctx, Message: msg, }); result.Error != nil { c.logger.Warn(c.ctx, "message handle canceled at txn", mlog.Err(result.Error)) @@ -265,3 +269,19 @@ func (c *consumerImpl) handleTxnMessage(msg message.ImmutableMessage) error { } return nil } + +func startDistConsumeSpanForMessage(ctx context.Context, msg message.ImmutableMessage) context.Context { + ctx, span := message.StartSpanForMessage(ctx, msg, message.SpanNameWALDistConsume) + message.OverwriteTraceContext(ctx, msg) + span.End() + return ctx +} + +func overwriteTxnMessagesTraceContext(ctx context.Context, txnMsg message.ImmutableTxnMessage) { + message.OverwriteTraceContext(ctx, txnMsg.Begin()) + _ = txnMsg.RangeOver(func(msg message.ImmutableMessage) error { + message.OverwriteTraceContext(ctx, msg) + return nil + }) + message.OverwriteTraceContext(ctx, txnMsg.Commit()) +} diff --git a/internal/streamingnode/client/handler/consumer/consumer_test.go b/internal/streamingnode/client/handler/consumer/consumer_test.go index 7dfcaa7055..d93406fadc 100644 --- a/internal/streamingnode/client/handler/consumer/consumer_test.go +++ b/internal/streamingnode/client/handler/consumer/consumer_test.go @@ -8,6 +8,11 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" + "go.opentelemetry.io/otel/trace" "github.com/milvus-io/milvus-proto/go-api/v3/msgpb" "github.com/milvus-io/milvus/pkg/v3/mocks/proto/mock_streamingpb" @@ -110,6 +115,202 @@ func TestConsumerWithCancellation(t *testing.T) { assert.ErrorIs(t, c.consumer.Error(), context.Canceled) } +func TestRemoteConsumerOverwritesTraceContextWithDistConsumeSpan(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), message.SpanNameWALCatchupConsume) + sourceSC := trace.SpanContextFromContext(sourceCtx) + sourceSpan.End() + + h := &captureTraceHandler{ch: make(chan message.HandleParam, 1)} + c := newMockedConsumerImpl(t, context.Background(), h) + + mmsg, _ := message.NewInsertMessageBuilderV1(). + WithHeader(&message.InsertMessageHeader{}). + WithBody(&msgpb.InsertRequest{}). + WithVChannel("test-1"). + BuildMutable() + message.InjectTraceContext(sourceCtx, mmsg) + c.recvCh <- newConsumeResponse(walimplstest.NewTestMessageID(1), mmsg) + + var param message.HandleParam + select { + case param = <-h.ch: + case <-time.After(time.Second): + t.Fatal("timeout waiting for consumed message") + } + + c.consumer.Close() + <-c.consumer.Done() + require.NoError(t, c.consumer.Error()) + + spans := exporter.GetSpans() + var distConsume tracetest.SpanStub + for _, s := range spans { + if s.Name == message.SpanNameWALDistConsume { + distConsume = s + break + } + } + require.Equal(t, message.SpanNameWALDistConsume, distConsume.Name) + assert.Equal(t, sourceSC.TraceID(), distConsume.SpanContext.TraceID()) + assert.Equal(t, sourceSC.SpanID(), distConsume.Parent.SpanID()) + + ctxSC := trace.SpanContextFromContext(param.Ctx) + assert.Equal(t, distConsume.SpanContext.TraceID(), ctxSC.TraceID()) + assert.Equal(t, distConsume.SpanContext.SpanID(), ctxSC.SpanID()) + + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), param.Message)) + assert.Equal(t, distConsume.SpanContext.TraceID(), msgSC.TraceID()) + assert.Equal(t, distConsume.SpanContext.SpanID(), msgSC.SpanID()) +} + +func TestRemoteConsumerSkipsTraceForTimeTickMessage(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), message.SpanNameWALCatchupConsume) + sourceSpan.End() + + h := &captureTraceHandler{ch: make(chan message.HandleParam, 1)} + c := newMockedConsumerImpl(t, context.Background(), h) + + msgID := walimplstest.NewTestMessageID(1) + mmsg := message.CreateTestTimeTickSyncMessage(t, 1, 100, msgID) + message.InjectTraceContext(sourceCtx, mmsg) + c.recvCh <- newConsumeResponse(msgID, mmsg) + + var param message.HandleParam + select { + case param = <-h.ch: + case <-time.After(time.Second): + t.Fatal("timeout waiting for consumed message") + } + + c.consumer.Close() + <-c.consumer.Done() + require.NoError(t, c.consumer.Error()) + + for _, s := range exporter.GetSpans() { + assert.NotEqual(t, message.SpanNameWALDistConsume, s.Name) + } + + ctxSC := trace.SpanContextFromContext(param.Ctx) + assert.False(t, ctxSC.IsValid()) + + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), param.Message)) + assert.False(t, msgSC.IsValid()) +} + +func TestRemoteConsumerStartsDistConsumeSpanOnlyOnTxnCommit(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), message.SpanNameWALCatchupConsume) + sourceSC := trace.SpanContextFromContext(sourceCtx) + sourceSpan.End() + + h := &captureTraceHandler{ch: make(chan message.HandleParam, 1)} + c := newMockedConsumerImpl(t, context.Background(), h) + + txnCtx := message.TxnContext{ + TxnID: 1, + Keepalive: time.Second, + } + begin := message.NewBeginTxnMessageBuilderV2(). + WithVChannel("test-1"). + WithHeader(&message.BeginTxnMessageHeader{}). + WithBody(&message.BeginTxnMessageBody{}). + MustBuildMutable(). + WithTxnContext(txnCtx) + message.InjectTraceContext(sourceCtx, begin) + c.recvCh <- newConsumeResponse(walimplstest.NewTestMessageID(1), begin) + + body := message.CreateTestEmptyInsertMesage(1, nil).WithTxnContext(txnCtx) + message.InjectTraceContext(sourceCtx, body) + c.recvCh <- newConsumeResponse(walimplstest.NewTestMessageID(2), body) + + commit := message.NewCommitTxnMessageBuilderV2(). + WithVChannel("test-1"). + WithHeader(&message.CommitTxnMessageHeader{}). + WithBody(&message.CommitTxnMessageBody{}). + MustBuildMutable(). + WithTxnContext(txnCtx) + message.InjectTraceContext(sourceCtx, commit) + c.recvCh <- newConsumeResponse(walimplstest.NewTestMessageID(3), commit) + + var param message.HandleParam + select { + case param = <-h.ch: + case <-time.After(time.Second): + t.Fatal("timeout waiting for consumed txn message") + } + + c.consumer.Close() + <-c.consumer.Done() + require.NoError(t, c.consumer.Error()) + + distConsumes := findSpansByName(exporter.GetSpans(), message.SpanNameWALDistConsume) + require.Len(t, distConsumes, 1) + assert.Equal(t, sourceSC.TraceID(), distConsumes[0].SpanContext.TraceID()) + assert.Equal(t, sourceSC.SpanID(), distConsumes[0].Parent.SpanID()) + assert.Equal(t, message.MessageTypeTxn, param.Message.MessageType()) + + ctxSC := trace.SpanContextFromContext(param.Ctx) + assert.Equal(t, distConsumes[0].SpanContext.TraceID(), ctxSC.TraceID()) + assert.Equal(t, distConsumes[0].SpanContext.SpanID(), ctxSC.SpanID()) + + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), param.Message)) + assert.Equal(t, distConsumes[0].SpanContext.TraceID(), msgSC.TraceID()) + assert.Equal(t, distConsumes[0].SpanContext.SpanID(), msgSC.SpanID()) + + txnMsg := message.AsImmutableTxnMessage(param.Message) + require.NotNil(t, txnMsg) + assertMessageTraceContext(t, txnMsg.Begin(), distConsumes[0].SpanContext) + assert.NoError(t, txnMsg.RangeOver(func(msg message.ImmutableMessage) error { + assertMessageTraceContext(t, msg, distConsumes[0].SpanContext) + return nil + })) + assertMessageTraceContext(t, txnMsg.Commit(), distConsumes[0].SpanContext) +} + +func findSpansByName(spans tracetest.SpanStubs, name string) []tracetest.SpanStub { + result := make([]tracetest.SpanStub, 0) + for _, s := range spans { + if s.Name == name { + result = append(result, s) + } + } + return result +} + +func assertMessageTraceContext(t *testing.T, msg message.ImmutableMessage, expected trace.SpanContext) { + t.Helper() + + sc := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), msg)) + assert.Equal(t, expected.TraceID(), sc.TraceID()) + assert.Equal(t, expected.SpanID(), sc.SpanID()) +} + type mockedConsumer struct { consumer Consumer recvCh chan *streamingpb.ConsumeResponse @@ -188,3 +389,15 @@ func newConsumeResponse(id message.MessageID, msg message.MutableMessage) *strea }, } } + +type captureTraceHandler struct { + ch chan message.HandleParam +} + +func (h *captureTraceHandler) Handle(param message.HandleParam) message.HandleResult { + h.ch <- param + return message.HandleResult{MessageHandled: true} +} + +func (h *captureTraceHandler) Close() { +} diff --git a/internal/streamingnode/client/handler/producer/produce_grpc_client.go b/internal/streamingnode/client/handler/producer/produce_grpc_client.go index 3b3d26b33c..9c52cdf382 100644 --- a/internal/streamingnode/client/handler/producer/produce_grpc_client.go +++ b/internal/streamingnode/client/handler/producer/produce_grpc_client.go @@ -1,6 +1,8 @@ package producer import ( + "context" + "github.com/milvus-io/milvus/pkg/v3/proto/streamingpb" "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" ) @@ -11,7 +13,8 @@ type produceGrpcClient struct { } // SendProduceMessage sends the produce message to server. -func (p *produceGrpcClient) SendProduceMessage(requestID int64, msg message.MutableMessage) error { +func (p *produceGrpcClient) SendProduceMessage(ctx context.Context, requestID int64, msg message.MutableMessage) error { + message.OverwriteTraceContext(ctx, msg) return p.Send(&streamingpb.ProduceRequest{ Request: &streamingpb.ProduceRequest_Produce{ Produce: &streamingpb.ProduceMessageRequest{ diff --git a/internal/streamingnode/client/handler/producer/producer_impl.go b/internal/streamingnode/client/handler/producer/producer_impl.go index a67bf5627a..a3ffc21f22 100644 --- a/internal/streamingnode/client/handler/producer/producer_impl.go +++ b/internal/streamingnode/client/handler/producer/producer_impl.go @@ -259,7 +259,7 @@ func (p *producerImpl) sendLoop() (err error) { // Store the request to pending request map. p.pendingRequests.Store(requestID, req) // Send the produce message to server. - if err := p.grpcStreamClient.SendProduceMessage(requestID, req.msg); err != nil { + if err := p.grpcStreamClient.SendProduceMessage(req.ctx, requestID, req.msg); err != nil { // If send failed, remove the request from pending request map and return error to client. p.notifyRequest(requestID, produceResponse{ err: err, diff --git a/internal/streamingnode/client/handler/producer/producer_test.go b/internal/streamingnode/client/handler/producer/producer_test.go index 9c3837d526..b5a37900b2 100644 --- a/internal/streamingnode/client/handler/producer/producer_test.go +++ b/internal/streamingnode/client/handler/producer/producer_test.go @@ -8,6 +8,11 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" + "go.opentelemetry.io/otel/trace" "github.com/milvus-io/milvus/internal/util/streamingutil/status" "github.com/milvus-io/milvus/pkg/v3/mocks/proto/mock_streamingpb" @@ -156,3 +161,94 @@ func TestProducer(t *testing.T) { producer.Close() assert.False(t, producer.IsAvailable()) } + +func TestProducerAppendOverwritesTraceContextDuringSerialization(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + c := mock_streamingpb.NewMockStreamingNodeHandlerServiceClient(t) + cc := mock_streamingpb.NewMockStreamingNodeHandlerService_ProduceClient(t) + recvCh := make(chan *streamingpb.ProduceResponse, 10) + cc.EXPECT().Recv().RunAndReturn(func() (*streamingpb.ProduceResponse, error) { + msg, ok := <-recvCh + if !ok { + return nil, io.EOF + } + return msg, nil + }) + sendCh := make(chan *streamingpb.ProduceRequest, 1) + cc.EXPECT().Send(mock.Anything).RunAndReturn(func(pr *streamingpb.ProduceRequest) error { + sendCh <- pr + return nil + }) + c.EXPECT().Produce(mock.Anything, mock.Anything).Return(cc, nil) + cc.EXPECT().CloseSend().RunAndReturn(func() error { + recvCh <- &streamingpb.ProduceResponse{Response: &streamingpb.ProduceResponse_Close{}} + close(recvCh) + return nil + }) + + opts := &ProducerOptions{ + Assignment: &types.PChannelInfoAssigned{ + Channel: types.PChannelInfo{Name: "test", Term: 1}, + Node: types.StreamingNodeInfo{ServerID: 1, Address: "localhost"}, + }, + } + + recvCh <- &streamingpb.ProduceResponse{ + Response: &streamingpb.ProduceResponse_Create{ + Create: &streamingpb.CreateProducerResponse{}, + }, + } + producer, err := CreateProducer(context.Background(), opts, c) + require.NoError(t, err) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), "source") + sourceSpan.End() + distCtx, distSpan := otel.Tracer("test").Start(sourceCtx, message.SpanNameWALDistAppend) + distSC := trace.SpanContextFromContext(distCtx) + + msg := message.CreateTestEmptyInsertMesage(1, nil) + message.InjectTraceContext(sourceCtx, msg) + appendDone := make(chan struct{}) + go func() { + defer close(appendDone) + result, err := producer.Append(distCtx, msg) + assert.NoError(t, err) + assert.NotNil(t, result) + }() + + req := <-sendCh + serializedMsg := req.GetProduce().GetMessage() + serializedImmutable := message.NewImmutableMesasge( + walimplstest.NewTestMessageID(1), + serializedMsg.GetPayload(), + serializedMsg.GetProperties(), + ) + serializedSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), serializedImmutable)) + assert.Equal(t, distSC.TraceID(), serializedSC.TraceID()) + assert.Equal(t, distSC.SpanID(), serializedSC.SpanID()) + + recvCh <- &streamingpb.ProduceResponse{ + Response: &streamingpb.ProduceResponse_Produce{ + Produce: &streamingpb.ProduceMessageResponse{ + RequestId: req.GetProduce().GetRequestId(), + Response: &streamingpb.ProduceMessageResponse_Result{ + Result: &streamingpb.ProduceMessageResponseResult{ + Id: walimplstest.NewTestMessageID(1).IntoProto(), + LastConfirmedId: walimplstest.NewTestMessageID(1).IntoProto(), + }, + }, + }, + }, + } + <-appendDone + distSpan.End() + producer.Close() +} diff --git a/internal/streamingnode/server/flusher/flusherimpl/flusher_components.go b/internal/streamingnode/server/flusher/flusherimpl/flusher_components.go index 1fae4bf0e0..13868890e3 100644 --- a/internal/streamingnode/server/flusher/flusherimpl/flusher_components.go +++ b/internal/streamingnode/server/flusher/flusherimpl/flusher_components.go @@ -38,11 +38,11 @@ type flusherComponents struct { } // WhenCreateCollection handles the create collection message. -func (impl *flusherComponents) WhenCreateCollection(createCollectionMsg message.ImmutableCreateCollectionMessageV1) { +func (impl *flusherComponents) WhenCreateCollection(ctx context.Context, createCollectionMsg message.ImmutableCreateCollectionMessageV1) { // because we need to get the schema from the recovery storage, we need to observe the message at recovery storage first. - impl.rs.ObserveMessage(context.Background(), createCollectionMsg) + impl.rs.ObserveMessage(ctx, createCollectionMsg) if _, ok := impl.dataServices[createCollectionMsg.VChannel()]; ok { - impl.logger.Info(context.TODO(), "the data sync service of current vchannel is built, skip it", mlog.FieldVChannel(createCollectionMsg.VChannel())) + impl.logger.Info(ctx, "the data sync service of current vchannel is built, skip it", mlog.FieldVChannel(createCollectionMsg.VChannel())) // May repeated consumed, so we ignore the message. return } @@ -50,7 +50,7 @@ func (impl *flusherComponents) WhenCreateCollection(createCollectionMsg message. // It should already be recovered from the recovery storage. // if it's not in recovery storage, it means the createCollection is already dropped. // so we can skip it. - impl.logger.Info(context.TODO(), "the create collection message is older than the recovery checkpoint, skip it", + impl.logger.Info(ctx, "the create collection message is older than the recovery checkpoint, skip it", mlog.FieldVChannel(createCollectionMsg.VChannel()), mlog.Uint64("timeTick", createCollectionMsg.TimeTick()), mlog.Uint64("recoveryCheckPointTimeTick", impl.recoveryCheckPointTimeTick)) @@ -100,16 +100,16 @@ func (impl *flusherComponents) WhenCreateCollection(createCollectionMsg message. }, nil, ) - impl.addNewDataSyncService(createCollectionMsg, msgChan, ds) + impl.addNewDataSyncService(ctx, createCollectionMsg, msgChan, ds) } // WhenDropCollection handles the drop collection message. -func (impl *flusherComponents) WhenDropCollection(vchannel string) { +func (impl *flusherComponents) WhenDropCollection(ctx context.Context, vchannel string) { // flowgraph is removed by data sync service it self. if ds, ok := impl.dataServices[vchannel]; ok { ds.Close() delete(impl.dataServices, vchannel) - impl.logger.Info(context.TODO(), "drop data sync service", mlog.FieldVChannel(vchannel)) + impl.logger.Info(ctx, "drop data sync service", mlog.FieldVChannel(vchannel)) } } @@ -141,6 +141,7 @@ func (impl *flusherComponents) broadcastToAllDataSyncService(ctx context.Context // addNewDataSyncService adds a new data sync service to the components when new collection is created. func (impl *flusherComponents) addNewDataSyncService( + ctx context.Context, createCollectionMsg message.ImmutableCreateCollectionMessageV1, input chan<- *msgstream.MsgPack, ds *pipeline.DataSyncService, @@ -148,7 +149,7 @@ func (impl *flusherComponents) addNewDataSyncService( newDS := newDataSyncServiceWrapper(createCollectionMsg.VChannel(), input, ds, createCollectionMsg.TimeTick()) newDS.Start() impl.dataServices[createCollectionMsg.VChannel()] = newDS - impl.logger.Info(context.TODO(), "create data sync service done", mlog.FieldVChannel(createCollectionMsg.VChannel())) + impl.logger.Info(ctx, "create data sync service done", mlog.FieldVChannel(createCollectionMsg.VChannel())) } // Close release all the resources of components. diff --git a/internal/streamingnode/server/flusher/flusherimpl/wal_flusher.go b/internal/streamingnode/server/flusher/flusherimpl/wal_flusher.go index 8c720e4821..3761968724 100644 --- a/internal/streamingnode/server/flusher/flusherimpl/wal_flusher.go +++ b/internal/streamingnode/server/flusher/flusherimpl/wal_flusher.go @@ -232,6 +232,7 @@ func (impl *WALFlusherImpl) generateScanner(ctx context.Context, l wal.WAL, chec // dispatch dispatches the message to the related handler for flusher components. func (impl *WALFlusherImpl) dispatch(msg message.ImmutableMessage) (err error) { + ctx := message.ExtractTraceContext(impl.notifier.Context(), msg) if msg.MessageType() == message.MessageTypeTimeTick && !msg.IsPersisted() { // Currently, milvus use the timetick to synchronize the system periodically, // so the wal will still produce empty timetick message after the last write operation is done. @@ -253,7 +254,7 @@ func (impl *WALFlusherImpl) dispatch(msg message.ImmutableMessage) (err error) { if msg.MessageType() == message.MessageTypeCommitImport { // CommitImport must not be observed until DataCoord accepts the commit // fence; otherwise replay can skip the only retry signal for this vchannel. - return impl.dispatchCommitImport(msg) + return impl.dispatchCommitImport(ctx, msg) } // TODO: should be removed at 3.0, after merge the flusher logic into recovery storage. @@ -261,16 +262,16 @@ func (impl *WALFlusherImpl) dispatch(msg message.ImmutableMessage) (err error) { // Other messages should keep the deferred order so lifecycle cleanup such as // DropCollection can finish the flowgraph before recovery storage observes it. if msg.MessageType() == message.MessageTypeTruncateCollection { - if err := impl.ObserveMessage(impl.notifier.Context(), msg); err != nil { - impl.logger.Warn(context.TODO(), "failed to observe message", mlog.Err(err)) + if err := impl.ObserveMessage(ctx, msg); err != nil { + impl.logger.Warn(ctx, "failed to observe message", mlog.Err(err)) return err } } else { // TODO: We will merge the flusher into recovery storage in future. // Currently, flusher works as a separate component. defer func() { - if err = impl.ObserveMessage(impl.notifier.Context(), msg); err != nil { - impl.logger.Warn(context.TODO(), "failed to observe message", mlog.Err(err)) + if err = impl.ObserveMessage(ctx, msg); err != nil { + impl.logger.Warn(ctx, "failed to observe message", mlog.Err(err)) } }() } @@ -285,33 +286,33 @@ func (impl *WALFlusherImpl) dispatch(msg message.ImmutableMessage) (err error) { case message.MessageTypeCreateCollection: createCollectionMsg, err := message.AsImmutableCreateCollectionMessageV1(msg) if err != nil { - impl.logger.DPanic(context.TODO(), "the message type is not CreateCollectionMessage", mlog.Err(err)) + impl.logger.DPanic(ctx, "the message type is not CreateCollectionMessage", mlog.Err(err)) return nil } - impl.flusherComponents.WhenCreateCollection(createCollectionMsg) + impl.flusherComponents.WhenCreateCollection(ctx, createCollectionMsg) case message.MessageTypeDropCollection: // defer to remove the data sync service from the components. // TODO: Current drop collection message will be handled by the underlying data sync service. defer func() { - impl.flusherComponents.WhenDropCollection(msg.VChannel()) + impl.flusherComponents.WhenDropCollection(ctx, msg.VChannel()) }() case message.MessageTypeRollbackImport: // No-op: DataCoord DDL ack callback handles all state changes. - impl.logger.Info(context.TODO(), "RollbackImportMessage consumed (no-op in flusher)", + impl.logger.Info(ctx, "RollbackImportMessage consumed (no-op in flusher)", mlog.FieldVChannel(msg.VChannel())) return nil // don't forward to flusherComponents } - return impl.flusherComponents.HandleMessage(impl.notifier.Context(), msg) + return impl.flusherComponents.HandleMessage(ctx, msg) } -func (impl *WALFlusherImpl) dispatchCommitImport(msg message.ImmutableMessage) error { +func (impl *WALFlusherImpl) dispatchCommitImport(ctx context.Context, msg message.ImmutableMessage) error { if funcutil.IsControlChannel(msg.VChannel()) && !msg.IsPChannelLevel() { - return impl.ObserveMessage(impl.notifier.Context(), msg) + return impl.ObserveMessage(ctx, msg) } commitMsg, err := message.AsImmutableCommitImportMessageV2(msg) if err != nil { - impl.logger.DPanic(context.TODO(), "failed to parse CommitImportMessage", mlog.Err(err)) + impl.logger.DPanic(ctx, "failed to parse CommitImportMessage", mlog.Err(err)) return nil } vchannel := msg.VChannel() @@ -319,17 +320,17 @@ func (impl *WALFlusherImpl) dispatchCommitImport(msg message.ImmutableMessage) e // Flush DML data before this commit fence. Panic on failure so WAL replays the message. if err := resource.Resource().WriteBufferManager(). - FlushChannel(context.Background(), vchannel, msg.TimeTick()); err != nil { + FlushChannel(ctx, vchannel, msg.TimeTick()); err != nil { if errors.Is(err, merr.ErrChannelNotFound) { - impl.logger.Info(context.TODO(), "CommitImport targets stale vchannel, skip local flush and continue commit ack", + impl.logger.Info(ctx, "CommitImport targets stale vchannel, skip local flush and continue commit ack", mlog.FieldVChannel(vchannel), mlog.FieldJobID(jobID), mlog.Err(err)) } else { - impl.logger.Panic(context.TODO(), "FlushChannel on CommitImport failed, panicking to retry from WAL", + impl.logger.Panic(ctx, "FlushChannel on CommitImport failed, panicking to retry from WAL", mlog.FieldVChannel(vchannel), mlog.FieldJobID(jobID), mlog.Err(err)) } } - mixCoord, err := resource.Resource().MixCoordClient().GetWithContext(impl.notifier.Context()) + mixCoord, err := resource.Resource().MixCoordClient().GetWithContext(ctx) if err != nil { return errors.Wrap(err, "failed to get MixCoordClient for HandleCommitVchannel") } @@ -337,19 +338,19 @@ func (impl *WALFlusherImpl) dispatchCommitImport(msg message.ImmutableMessage) e // fence for this dispatch and must not be repeated on every retry. // This retry blocks the whole pchannel flusher until DataCoord accepts the // commit fence, preserving WAL replay order for later messages on the pchannel. - impl.logger.Info(impl.notifier.Context(), "HandleCommitVchannel waits until DataCoord accepts the commit fence", + impl.logger.Info(ctx, "HandleCommitVchannel waits until DataCoord accepts the commit fence", mlog.FieldJobID(jobID), mlog.FieldVChannel(vchannel), mlog.Uint64("commitTs", msg.TimeTick())) - if err := retry.Do(impl.notifier.Context(), func() error { - resp, err := mixCoord.HandleCommitVchannel(impl.notifier.Context(), &datapb.HandleCommitVchannelRequest{ + if err := retry.Do(ctx, func() error { + resp, err := mixCoord.HandleCommitVchannel(ctx, &datapb.HandleCommitVchannelRequest{ Base: commonpbutil.NewMsgBase(commonpbutil.WithSourceID(paramtable.GetNodeID())), JobId: jobID, Vchannel: vchannel, CommitTimestamp: msg.TimeTick(), }) if err := merr.CheckRPCCall(resp, err); err != nil { - impl.logger.Debug(context.TODO(), "HandleCommitVchannel failed, retry later", + impl.logger.Debug(ctx, "HandleCommitVchannel failed, retry later", mlog.FieldJobID(jobID), mlog.FieldVChannel(vchannel), mlog.Uint64("commitTs", msg.TimeTick()), @@ -361,13 +362,13 @@ func (impl *WALFlusherImpl) dispatchCommitImport(msg message.ImmutableMessage) e return err } - if err := impl.ObserveMessage(impl.notifier.Context(), msg); err != nil { - impl.logger.Warn(context.TODO(), "failed to observe CommitImport message", + if err := impl.ObserveMessage(ctx, msg); err != nil { + impl.logger.Warn(ctx, "failed to observe CommitImport message", mlog.FieldVChannel(vchannel), mlog.FieldJobID(jobID), mlog.Err(err)) return err } - impl.logger.Info(context.TODO(), "CommitImportMessage handled: vchannel committed", + impl.logger.Info(ctx, "CommitImportMessage handled: vchannel committed", mlog.FieldVChannel(vchannel), mlog.FieldJobID(jobID)) return nil } diff --git a/internal/streamingnode/server/flusher/flusherimpl/wal_flusher_test.go b/internal/streamingnode/server/flusher/flusherimpl/wal_flusher_test.go index ba1ca339af..57e7d021f5 100644 --- a/internal/streamingnode/server/flusher/flusherimpl/wal_flusher_test.go +++ b/internal/streamingnode/server/flusher/flusherimpl/wal_flusher_test.go @@ -13,6 +13,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/trace" "google.golang.org/grpc" "github.com/milvus-io/milvus-proto/go-api/v3/commonpb" @@ -168,6 +169,44 @@ func TestWALFlusher_DispatchObservesTruncateCollectionBeforeHandlingWithoutAckSy require.ErrorContains(t, flusher.dispatch(msg), "observe failed") } +func TestWALFlusherDispatchRestoresTraceContext(t *testing.T) { + expectedTraceID, err := trace.TraceIDFromHex("0102030405060708090a0b0c0d0e0f10") + require.NoError(t, err) + spanID, err := trace.SpanIDFromHex("0102030405060708") + require.NoError(t, err) + clientCtx := trace.ContextWithSpanContext(context.Background(), trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: expectedTraceID, + SpanID: spanID, + })) + + mutableMsg := message.NewDropCollectionMessageBuilderV1(). + WithHeader(&message.DropCollectionMessageHeader{ + CollectionId: 100, + }). + WithBody(&msgpb.DropCollectionRequest{ + Base: &commonpb.MsgBase{}, + }). + WithVChannel("vchannel-1"). + MustBuildMutable(). + WithTimeTick(100). + WithLastConfirmed(rmq.NewRmqID(1)) + message.InjectTraceContext(clientCtx, mutableMsg) + msg := mutableMsg.IntoImmutableMessage(rmq.NewRmqID(2)) + + var observedTraceID trace.TraceID + rs := mock_recovery.NewMockRecoveryStorage(t) + rs.EXPECT().ObserveMessage(mock.Anything, mock.Anything). + RunAndReturn(func(ctx context.Context, msg message.ImmutableMessage) error { + observedTraceID = trace.SpanContextFromContext(ctx).TraceID() + return nil + }). + Once() + + flusher := newTestWALFlusher(rs) + require.NoError(t, flusher.dispatch(msg)) + assert.Equal(t, expectedTraceID, observedTraceID) +} + func newTestWALFlusher(rs recovery.RecoveryStorage) *WALFlusherImpl { return &WALFlusherImpl{ notifier: syncutil.NewAsyncTaskNotifier[struct{}](), diff --git a/internal/streamingnode/server/service/handler/producer/produce_server.go b/internal/streamingnode/server/service/handler/producer/produce_server.go index 4313d9ff2d..b6a2556c25 100644 --- a/internal/streamingnode/server/service/handler/producer/produce_server.go +++ b/internal/streamingnode/server/service/handler/producer/produce_server.go @@ -188,13 +188,14 @@ func (p *ProduceServer) handleProduce(req *streamingpb.ProduceMessageRequest) { } p.appendWG.Add(1) - p.logger.Debug(context.TODO(), "recv produce message from client", mlog.Int64("requestID", req.RequestId)) - // Update metrics. msg := message.NewMutableMessageBeforeAppend(req.GetMessage().GetPayload(), req.GetMessage().GetProperties()) + ctx := message.ExtractTraceContext(p.produceServer.Context(), msg) + p.logger.Debug(ctx, "recv produce message from client", mlog.Int64("requestID", req.RequestId)) + // Update metrics. metricsGuard := p.metrics.StartProduce() if err := p.validateMessage(msg); err != nil { - p.logger.Warn(context.TODO(), "produce message validation failed", mlog.Int64("requestID", req.RequestId), mlog.Err(err)) - p.sendProduceResult(req.RequestId, nil, err) + p.logger.Warn(ctx, "produce message validation failed", mlog.Int64("requestID", req.RequestId), mlog.Err(err)) + p.sendProduceResult(ctx, req.RequestId, nil, err) metricsGuard.Finish(err) p.appendWG.Done() return @@ -202,12 +203,12 @@ func (p *ProduceServer) handleProduce(req *streamingpb.ProduceMessageRequest) { // Append message to wal. // Concurrent append request can be executed concurrently. - p.wal.AppendAsync(p.produceServer.Context(), msg, func(appendResult *wal.AppendResult, err error) { + p.wal.AppendAsync(ctx, msg, func(appendResult *wal.AppendResult, err error) { defer func() { metricsGuard.Finish(err) p.appendWG.Done() }() - p.sendProduceResult(req.RequestId, appendResult, err) + p.sendProduceResult(ctx, req.RequestId, appendResult, err) }) } @@ -255,12 +256,12 @@ func (p *ProduceServer) UpdateRateLimitState(state ratelimit.RateLimitState) { } // sendProduceResult sends the produce result to client. -func (p *ProduceServer) sendProduceResult(reqID int64, appendResult *wal.AppendResult, err error) { +func (p *ProduceServer) sendProduceResult(ctx context.Context, reqID int64, appendResult *wal.AppendResult, err error) { resp := &streamingpb.ProduceMessageResponse{ RequestId: reqID, } if err != nil { - p.logger.Warn(context.TODO(), "append message to wal failed", mlog.Int64("requestID", reqID), mlog.Err(err)) + p.logger.Warn(ctx, "append message to wal failed", mlog.Int64("requestID", reqID), mlog.Err(err)) resp.Response = &streamingpb.ProduceMessageResponse_Error{Error: status.AsStreamingError(err).AsPBError()} } else { resp.Response = &streamingpb.ProduceMessageResponse_Result{Result: appendResult.IntoProto()} @@ -270,9 +271,9 @@ func (p *ProduceServer) sendProduceResult(reqID int64, appendResult *wal.AppendR // all pending response message should be dropped, client side will handle it. select { case p.produceMessageCh <- resp: - p.logger.Debug(context.TODO(), "send produce message response to client", mlog.Int64("requestID", reqID), mlog.Any("appendResult", appendResult), mlog.Err(err)) + p.logger.Debug(ctx, "send produce message response to client", mlog.Int64("requestID", reqID), mlog.Any("appendResult", appendResult), mlog.Err(err)) case <-p.produceServer.Context().Done(): - p.logger.Warn(context.TODO(), "stream closed before produce message response sent", mlog.Int64("requestID", reqID), mlog.Any("appendResult", appendResult), mlog.Err(err)) + p.logger.Warn(ctx, "stream closed before produce message response sent", mlog.Int64("requestID", reqID), mlog.Any("appendResult", appendResult), mlog.Err(err)) return } } diff --git a/internal/streamingnode/server/service/handler/producer/produce_server_test.go b/internal/streamingnode/server/service/handler/producer/produce_server_test.go index b4ad090a70..fc51a3b160 100644 --- a/internal/streamingnode/server/service/handler/producer/produce_server_test.go +++ b/internal/streamingnode/server/service/handler/producer/produce_server_test.go @@ -913,7 +913,7 @@ func TestProduceServerSendProduceResult_ContextCanceled(t *testing.T) { // This should not block and should log warning msgID := walimplstest.NewTestMessageID(1) - p.sendProduceResult(1, &wal.AppendResult{ + p.sendProduceResult(context.Background(), 1, &wal.AppendResult{ MessageID: msgID, LastConfirmedMessageID: msgID, TimeTick: 100, diff --git a/internal/streamingnode/server/service/handler/producer/produce_server_trace_test.go b/internal/streamingnode/server/service/handler/producer/produce_server_trace_test.go new file mode 100644 index 0000000000..61725a88a9 --- /dev/null +++ b/internal/streamingnode/server/service/handler/producer/produce_server_trace_test.go @@ -0,0 +1,104 @@ +//go:build test && dynamic + +package producer + +import ( + "context" + "strconv" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "go.opentelemetry.io/otel/trace" + + "github.com/milvus-io/milvus/internal/mocks/streamingnode/server/mock_wal" + "github.com/milvus-io/milvus/internal/streamingnode/server/wal" + "github.com/milvus-io/milvus/pkg/v3/mlog" + "github.com/milvus-io/milvus/pkg/v3/mocks/proto/mock_streamingpb" + "github.com/milvus-io/milvus/pkg/v3/proto/messagespb" + "github.com/milvus-io/milvus/pkg/v3/proto/streamingpb" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/types" + "github.com/milvus-io/milvus/pkg/v3/streaming/walimpls/impls/walimplstest" +) + +// TestHandleProduce_ExtractsTraceContext verifies that handleProduce restores +// the client-injected trace context from message properties before appending. +func TestHandleProduce_ExtractsTraceContext(t *testing.T) { + expectedTraceID := trace.TraceID{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16} + clientCtx := trace.ContextWithSpanContext(context.Background(), trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: expectedTraceID, + SpanID: trace.SpanID{1, 2, 3, 4, 5, 6, 7, 8}, + TraceFlags: trace.FlagsSampled, + })) + + // Build properties that carry the injected trace context (simulating what the + // gRPC client does in Task 4). + props := map[string]string{ + "_v": "1", + "_t": strconv.FormatInt(int64(message.MessageTypeInsert), 10), + } + injectedMsg := message.NewMutableMessageBeforeAppend([]byte("test-payload"), props) + message.InjectTraceContext(clientCtx, injectedMsg) + + req := &streamingpb.ProduceMessageRequest{ + RequestId: 42, + Message: &messagespb.Message{ + Payload: []byte("test-payload"), + Properties: injectedMsg.Properties().ToRawMap(), + }, + } + + // Build a mock WAL; capture the ctx passed to AppendAsync. + l := mock_wal.NewMockWAL(t) + l.EXPECT().Channel().Return(types.PChannelInfo{Name: "test-ch", Term: 1}).Maybe() + l.EXPECT().IsAvailable().Return(true) + + var ( + capturedCtx context.Context + wg sync.WaitGroup + ) + wg.Add(1) + + l.EXPECT().AppendAsync(mock.Anything, mock.Anything, mock.Anything). + Run(func(ctx context.Context, msg message.MutableMessage, cb func(*wal.AppendResult, error)) { + capturedCtx = ctx + msgID := walimplstest.NewTestMessageID(1) + cb(&wal.AppendResult{ + MessageID: msgID, + LastConfirmedMessageID: msgID, + TimeTick: 100, + }, nil) + wg.Done() + }).Return() + + // Build a minimal ProduceServer (same pattern as existing tests). + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + grpcProduceServer := mock_streamingpb.NewMockStreamingNodeHandlerService_ProduceServer(t) + grpcProduceServer.EXPECT().Context().Return(ctx).Maybe() + grpcProduceServer.EXPECT().Send(mock.Anything).Return(nil).Maybe() + + ps := &ProduceServer{ + wal: l, + produceServer: &produceGrpcServerHelper{ + StreamingNodeHandlerService_ProduceServer: grpcProduceServer, + }, + logger: mlog.With(), + produceMessageCh: make(chan *streamingpb.ProduceMessageResponse, 10), + appendWG: sync.WaitGroup{}, + metrics: newProducerMetrics(types.PChannelInfo{Name: "test-ch", Term: 1}), + } + + // Invoke handleProduce synchronously; the callback fires inside the mock. + ps.handleProduce(req) + wg.Wait() + + // The ctx passed to AppendAsync must carry the extracted remote span context. + assert.NotNil(t, capturedCtx) + capturedSC := trace.SpanContextFromContext(capturedCtx) + assert.True(t, capturedSC.IsValid(), "ctx passed to AppendAsync should carry a valid span context") + assert.Equal(t, expectedTraceID, capturedSC.TraceID(), "AppendAsync ctx must share the client trace ID") +} diff --git a/internal/streamingnode/server/wal/adaptor/scanner_switchable.go b/internal/streamingnode/server/wal/adaptor/scanner_switchable.go index d9ce2f1fc7..4a0a70b8d6 100644 --- a/internal/streamingnode/server/wal/adaptor/scanner_switchable.go +++ b/internal/streamingnode/server/wal/adaptor/scanner_switchable.go @@ -176,6 +176,9 @@ func (s *catchupScanner) consumeWithScanner(ctx context.Context, scanner walimpl // when we switch from tailing mode to catchup mode. continue } + if shouldStartConsumeSpan(msg) { + startConsumeSpanForMessage(ctx, msg) + } if err := s.HandleMessage(ctx, msg); err != nil { return nil, err } @@ -267,6 +270,10 @@ func (s *tailingScanner) Do(ctx context.Context) (switchableScanner, error) { if err != nil { return nil, err } + // Do not start wal.catchup_consume or overwrite _tc in tailing mode. + // WriteAheadBuffer readers share the same immutable message instance, + // including its properties map, across all tailing consumers on this + // pchannel. Mutating trace context here would race with other readers. if err := s.HandleMessage(ctx, tailingImmutableMesasge{msg}); err != nil { return nil, err } @@ -293,3 +300,17 @@ func isTailingScanImmutableMessage(msg message.ImmutableMessage) (message.Immuta } return msg, false } + +func shouldStartConsumeSpan(msg message.ImmutableMessage) bool { + if msg.TxnContext() == nil { + return true + } + return msg.MessageType() == message.MessageTypeCommitTxn +} + +func startConsumeSpanForMessage(ctx context.Context, msg message.ImmutableMessage) { + ctx = message.ExtractTraceContext(ctx, msg) + ctx, span := message.StartSpanForMessage(ctx, msg, message.SpanNameWALCatchupConsume) + message.OverwriteTraceContext(ctx, msg) + span.End() +} diff --git a/internal/streamingnode/server/wal/adaptor/wal_adaptor.go b/internal/streamingnode/server/wal/adaptor/wal_adaptor.go index ded35d29ea..117ec778b1 100644 --- a/internal/streamingnode/server/wal/adaptor/wal_adaptor.go +++ b/internal/streamingnode/server/wal/adaptor/wal_adaptor.go @@ -6,6 +6,7 @@ import ( "github.com/cenkalti/backoff/v4" "github.com/cockroachdb/errors" + "go.opentelemetry.io/otel/codes" "go.uber.org/atomic" "google.golang.org/protobuf/types/known/anypb" @@ -150,12 +151,21 @@ func (w *walAdaptorImpl) GetSalvageCheckpoint() []*utility.ReplicateCheckpoint { } // Append writes a record to the log. -func (w *walAdaptorImpl) Append(ctx context.Context, msg message.MutableMessage) (*wal.AppendResult, error) { +func (w *walAdaptorImpl) Append(ctx context.Context, msg message.MutableMessage) (_ *wal.AppendResult, err error) { if !w.lifetime.Add(typeutil.LifetimeStateWorking) { return nil, status.NewOnShutdownError("wal is on shutdown") } defer w.lifetime.Done() + ctx, span := message.StartSpanForMessage(ctx, msg, message.SpanNameWALAppend) + defer func() { + if err != nil { + span.RecordError(err) + span.SetStatus(codes.Error, err.Error()) + } + span.End() + }() + if w.isFenced.Load() { // if the wal is fenced, we should reject all append operations. return nil, status.NewChannelFenced(w.Channel().String()) @@ -206,7 +216,7 @@ func (w *walAdaptorImpl) Append(ctx context.Context, msg message.MutableMessage) }) metricsGuard.FinishAppend() if err != nil { - appendMetrics.Done(nil, err) + appendMetrics.Done(ctx, nil, err) if errors.Is(err, walimpls.ErrFenced) { // if the append operation of wal is fenced, we should report the error to the client. if w.isFenced.CompareAndSwap(false, true) { @@ -243,7 +253,7 @@ func (w *walAdaptorImpl) Append(ctx context.Context, msg message.MutableMessage) TxnCtx: extraAppendResult.TxnCtx, Extra: extra, } - appendMetrics.Done(r, nil) + appendMetrics.Done(ctx, r, nil) return r, nil } @@ -264,7 +274,14 @@ func (w *walAdaptorImpl) retryAppendWhenRecoverableError(ctx context.Context, ms // An append operation should be retried until it succeeds or some unrecoverable error occurs. for i := 0; ; i++ { - msgID, err := w.rwWALImpls.Append(ctx, msg) + appendCtx, span := message.StartSpanForMessage(ctx, msg, message.SpanNameWALAppendImpl) + message.OverwriteTraceContext(appendCtx, msg) + msgID, err := w.rwWALImpls.Append(appendCtx, msg) + if err != nil { + span.RecordError(err) + span.SetStatus(codes.Error, err.Error()) + } + span.End() if err == nil { if msg.MessageType() == message.MessageTypeAlterWAL { // if the append operation is a alter WAL message, we should log the message diff --git a/internal/streamingnode/server/wal/adaptor/wal_adaptor_trace_test.go b/internal/streamingnode/server/wal/adaptor/wal_adaptor_trace_test.go new file mode 100644 index 0000000000..2761479d6b --- /dev/null +++ b/internal/streamingnode/server/wal/adaptor/wal_adaptor_trace_test.go @@ -0,0 +1,372 @@ +package adaptor + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" + "go.opentelemetry.io/otel/trace" + + "github.com/milvus-io/milvus/internal/streamingnode/server/wal/metricsutil" + "github.com/milvus-io/milvus/internal/streamingnode/server/wal/utility" + "github.com/milvus-io/milvus/pkg/v3/mlog" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/types" + "github.com/milvus-io/milvus/pkg/v3/streaming/walimpls/helper" + "github.com/milvus-io/milvus/pkg/v3/streaming/walimpls/impls/walimplstest" +) + +func TestRetryAppendOverwritesTraceContextWithAppendImplSpan(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), "source") + sourceSpan.End() + + msg := message.CreateTestEmptyInsertMesage(1, nil) + message.InjectTraceContext(sourceCtx, msg) + + var capturedCtx context.Context + w := &walAdaptorImpl{ + rwWALImpls: newFirstTimeTickWALImpls(func(ctx context.Context, msg message.MutableMessage) (message.MessageID, error) { + capturedCtx = ctx + return walimplstest.NewTestMessageID(1), nil + }), + } + + _, err := w.retryAppendWhenRecoverableError(sourceCtx, msg) + require.NoError(t, err) + + spans := exporter.GetSpans() + var appendImpl tracetest.SpanStub + for _, s := range spans { + if s.Name == message.SpanNameWALAppendImpl { + appendImpl = s + break + } + } + require.Equal(t, message.SpanNameWALAppendImpl, appendImpl.Name) + + capturedSC := trace.SpanContextFromContext(capturedCtx) + assert.Equal(t, appendImpl.SpanContext.TraceID(), capturedSC.TraceID()) + assert.Equal(t, appendImpl.SpanContext.SpanID(), capturedSC.SpanID()) + + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), msg)) + assert.Equal(t, appendImpl.SpanContext.TraceID(), msgSC.TraceID()) + assert.Equal(t, appendImpl.SpanContext.SpanID(), msgSC.SpanID()) +} + +func TestRetryAppendSkipsTraceForTimeTickMessage(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), "source") + sourceSC := trace.SpanContextFromContext(sourceCtx) + sourceSpan.End() + + msgID := walimplstest.NewTestMessageID(1) + msg := message.CreateTestTimeTickSyncMessage(t, 1, 100, msgID) + message.InjectTraceContext(sourceCtx, msg) + + var capturedCtx context.Context + w := &walAdaptorImpl{ + rwWALImpls: newFirstTimeTickWALImpls(func(ctx context.Context, msg message.MutableMessage) (message.MessageID, error) { + capturedCtx = ctx + return msgID, nil + }), + } + + _, err := w.retryAppendWhenRecoverableError(sourceCtx, msg) + require.NoError(t, err) + + for _, s := range exporter.GetSpans() { + assert.NotEqual(t, message.SpanNameWALAppendImpl, s.Name) + } + + capturedSC := trace.SpanContextFromContext(capturedCtx) + assert.Equal(t, sourceSC.TraceID(), capturedSC.TraceID()) + assert.Equal(t, sourceSC.SpanID(), capturedSC.SpanID()) + + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), msg)) + assert.False(t, msgSC.IsValid()) +} + +func TestCatchupScannerOverwritesTraceContextWithConsumeSpan(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), "wal.appendimpl") + sourceSpan.End() + + msgID := walimplstest.NewTestMessageID(1) + mutableMsg := message.CreateTestEmptyInsertMesage(1, nil) + mutableMsg.WithTimeTick(100) + mutableMsg.WithLastConfirmed(msgID) + message.InjectTraceContext(sourceCtx, mutableMsg) + immutableMsg := mutableMsg.IntoImmutableMessage(msgID) + + capturedMsgs := runTraceTestCatchupScanner(t, immutableMsg) + require.Len(t, capturedMsgs, 1) + + spans := exporter.GetSpans() + var consume tracetest.SpanStub + for _, s := range spans { + if s.Name == message.SpanNameWALCatchupConsume { + consume = s + break + } + } + require.Equal(t, message.SpanNameWALCatchupConsume, consume.Name) + + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), capturedMsgs[0])) + assert.Equal(t, consume.SpanContext.TraceID(), msgSC.TraceID()) + assert.Equal(t, consume.SpanContext.SpanID(), msgSC.SpanID()) +} + +func TestCatchupScannerSkipsTraceForTimeTickMessage(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), "wal.appendimpl") + sourceSpan.End() + + msgID := walimplstest.NewTestMessageID(1) + mutableMsg := message.CreateTestTimeTickSyncMessage(t, 1, 100, msgID) + message.InjectTraceContext(sourceCtx, mutableMsg) + immutableMsg := mutableMsg.IntoImmutableMessage(msgID) + + capturedMsgs := runTraceTestCatchupScanner(t, immutableMsg) + require.Len(t, capturedMsgs, 1) + for _, s := range exporter.GetSpans() { + assert.NotEqual(t, message.SpanNameWALCatchupConsume, s.Name) + } + + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), capturedMsgs[0])) + assert.False(t, msgSC.IsValid()) +} + +func TestCatchupScannerStartsConsumeSpanOnlyOnTxnCommit(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), message.SpanNameWALTxn) + sourceSC := trace.SpanContextFromContext(sourceCtx) + sourceSpan.End() + + msgs := buildTraceTestTxnImmutableMessages(t, sourceCtx) + for _, msg := range msgs { + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), msg)) + assert.Equal(t, sourceSC.SpanID(), msgSC.SpanID()) + } + scanner := newTraceTestScannerAdaptor() + capturedMsgs := runTraceTestCatchupScanner(t, msgs...) + require.Len(t, capturedMsgs, 3) + for _, msg := range capturedMsgs { + scanner.handleUpstream(msg) + } + timeTickMsg := message.CreateTestTimeTickSyncMessage(t, 1, 102, walimplstest.NewTestMessageID(4)).IntoImmutableMessage(walimplstest.NewTestMessageID(4)) + scanner.handleUpstream(timeTickMsg) + + consumes := findTraceTestSpansByName(exporter.GetSpans(), message.SpanNameWALCatchupConsume) + require.Len(t, consumes, 1) + assert.Equal(t, sourceSC.TraceID(), consumes[0].SpanContext.TraceID()) + assert.Equal(t, sourceSC.SpanID(), consumes[0].Parent.SpanID()) + + capturedMsg := scanner.pendingQueue.Next() + require.Equal(t, message.MessageTypeTxn, capturedMsg.MessageType()) + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), capturedMsg)) + assert.Equal(t, consumes[0].SpanContext.TraceID(), msgSC.TraceID()) + assert.Equal(t, consumes[0].SpanContext.SpanID(), msgSC.SpanID()) +} + +func TestScannerAdaptorSkipsConsumeSpanForTailingMessage(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), message.SpanNameWALAppendImpl) + sourceSC := trace.SpanContextFromContext(sourceCtx) + sourceSpan.End() + + msgID := walimplstest.NewTestMessageID(1) + mutableMsg := message.CreateTestEmptyInsertMesage(1, nil) + mutableMsg.WithTimeTick(100) + mutableMsg.WithLastConfirmed(msgID) + message.InjectTraceContext(sourceCtx, mutableMsg) + immutableMsg := mutableMsg.IntoImmutableMessage(msgID) + + scanner := newTraceTestScannerAdaptor() + scanner.handleUpstream(tailingImmutableMesasge{immutableMsg}) + timeTickMsg := message.CreateTestTimeTickSyncMessage(t, 1, 100, msgID).IntoImmutableMessage(msgID) + scanner.handleUpstream(tailingImmutableMesasge{timeTickMsg}) + + for _, s := range exporter.GetSpans() { + assert.NotEqual(t, message.SpanNameWALCatchupConsume, s.Name) + } + + capturedMsg := scanner.pendingQueue.Next() + msgSC := trace.SpanContextFromContext(message.ExtractTraceContext(context.Background(), capturedMsg)) + assert.Equal(t, sourceSC.TraceID(), msgSC.TraceID()) + assert.Equal(t, sourceSC.SpanID(), msgSC.SpanID()) +} + +func newTraceTestScannerAdaptor() *scannerAdaptorImpl { + logger := mlog.With() + scanMetrics := metricsutil.NewScanMetrics(types.PChannelInfo{}).NewScannerMetrics() + return &scannerAdaptorImpl{ + logger: logger, + filterFunc: func(message.ImmutableMessage) bool { return true }, + reorderBuffer: utility.NewReOrderBuffer(), + pendingQueue: utility.NewPendingQueue(), + txnBuffer: utility.NewTxnBuffer(logger, scanMetrics), + ScannerHelper: helper.NewScannerHelper("trace-test"), + metrics: scanMetrics, + readRateCounter: utility.NewAverageRateCounter(10 * time.Second), + } +} + +func runTraceTestCatchupScanner(t *testing.T, msgs ...message.ImmutableMessage) []message.ImmutableMessage { + t.Helper() + + scannerCh := make(chan message.ImmutableMessage, len(msgs)) + for _, msg := range msgs { + scannerCh <- msg + } + close(scannerCh) + + msgCh := make(chan message.ImmutableMessage, len(msgs)) + scanner := &catchupScanner{ + switchableScannerImpl: switchableScannerImpl{ + scannerName: "trace-test", + logger: mlog.With(), + innerWAL: newFirstTimeTickWALImpls(nil), + msgChan: msgCh, + }, + } + _, err := scanner.consumeWithScanner(context.Background(), &traceTestScanner{ch: scannerCh}) + require.NoError(t, err) + + capturedMsgs := make([]message.ImmutableMessage, 0, len(msgs)) + for len(msgCh) > 0 { + capturedMsgs = append(capturedMsgs, <-msgCh) + } + return capturedMsgs +} + +func buildTraceTestTxnImmutableMessages(t *testing.T, ctx context.Context) []message.ImmutableMessage { + t.Helper() + + txnCtx := message.TxnContext{ + TxnID: 1, + Keepalive: time.Second, + } + lastConfirmed := walimplstest.NewTestMessageID(0) + + begin := message.NewBeginTxnMessageBuilderV2(). + WithVChannel("v1"). + WithHeader(&message.BeginTxnMessageHeader{}). + WithBody(&message.BeginTxnMessageBody{}). + MustBuildMutable(). + WithTxnContext(txnCtx). + WithTimeTick(100). + WithLastConfirmed(lastConfirmed) + message.InjectTraceContext(ctx, begin) + + body := message.CreateTestEmptyInsertMesage(1, nil). + WithTxnContext(txnCtx). + WithTimeTick(101). + WithLastConfirmed(lastConfirmed) + message.InjectTraceContext(ctx, body) + + commit := message.NewCommitTxnMessageBuilderV2(). + WithVChannel("v1"). + WithHeader(&message.CommitTxnMessageHeader{}). + WithBody(&message.CommitTxnMessageBody{}). + MustBuildMutable(). + WithTxnContext(txnCtx). + WithTimeTick(102). + WithLastConfirmed(lastConfirmed) + message.InjectTraceContext(ctx, commit) + + return []message.ImmutableMessage{ + begin.IntoImmutableMessage(walimplstest.NewTestMessageID(1)), + body.IntoImmutableMessage(walimplstest.NewTestMessageID(2)), + commit.IntoImmutableMessage(walimplstest.NewTestMessageID(3)), + } +} + +func findTraceTestSpansByName(spans tracetest.SpanStubs, name string) []tracetest.SpanStub { + result := make([]tracetest.SpanStub, 0) + for _, s := range spans { + if s.Name == name { + result = append(result, s) + } + } + return result +} + +type traceTestScanner struct { + ch <-chan message.ImmutableMessage +} + +func (s *traceTestScanner) Name() string { + return "trace-test-scanner" +} + +func (s *traceTestScanner) Chan() <-chan message.ImmutableMessage { + return s.ch +} + +func (s *traceTestScanner) Error() error { + return nil +} + +func (s *traceTestScanner) Done() <-chan struct{} { + done := make(chan struct{}) + close(done) + return done +} + +func (s *traceTestScanner) Close() error { + return nil +} diff --git a/internal/streamingnode/server/wal/metricsutil/append.go b/internal/streamingnode/server/wal/metricsutil/append.go index 0104c8f577..2ed66065aa 100644 --- a/internal/streamingnode/server/wal/metricsutil/append.go +++ b/internal/streamingnode/server/wal/metricsutil/append.go @@ -1,6 +1,7 @@ package metricsutil import ( + "context" "fmt" "time" @@ -128,10 +129,10 @@ func (m *AppendMetrics) RangeOverInterceptors(f func(name string, ims []*Interce } // Done push the metrics. -func (m *AppendMetrics) Done(result *types.AppendResult, err error) { +func (m *AppendMetrics) Done(ctx context.Context, result *types.AppendResult, err error) { m.err = err m.result = result - m.wm.done(m) + m.wm.done(ctx, m) } // InterceptorCollectGuard is used to collect the metrics of interceptor. diff --git a/internal/streamingnode/server/wal/metricsutil/wal_write.go b/internal/streamingnode/server/wal/metricsutil/wal_write.go index d75d790eb2..f91e5c6c36 100644 --- a/internal/streamingnode/server/wal/metricsutil/wal_write.go +++ b/internal/streamingnode/server/wal/metricsutil/wal_write.go @@ -74,7 +74,7 @@ func (m *WriteMetrics) StartAppend(msg message.MutableMessage) *AppendMetrics { } } -func (m *WriteMetrics) done(appendMetrics *AppendMetrics) { +func (m *WriteMetrics) done(ctx context.Context, appendMetrics *AppendMetrics) { if !appendMetrics.msg.IsPersisted() { return } @@ -96,17 +96,17 @@ func (m *WriteMetrics) done(appendMetrics *AppendMetrics) { } } if appendMetrics.err != nil { - m.Logger().Warn(context.TODO(), "append message into wal failed", appendMetrics.IntoLogFields()...) + m.Logger().Warn(ctx, "append message into wal failed", appendMetrics.IntoLogFields()...) return } if appendMetrics.appendDuration >= m.slowLogThreshold { // log slow append catch - m.Logger().Warn(context.TODO(), "append message into wal too slow", appendMetrics.IntoLogFields()...) + m.Logger().Warn(ctx, "append message into wal too slow", appendMetrics.IntoLogFields()...) return } logLV := appendMetrics.msg.MessageType().LogLevel() if m.Logger().LevelEnabled(logLV) { - m.Logger().Log(context.TODO(), logLV, "append message into wal", appendMetrics.IntoLogFields()...) + m.Logger().Log(ctx, logLV, "append message into wal", appendMetrics.IntoLogFields()...) } } diff --git a/internal/streamingnode/server/wal/metricsutil/wal_write_trace_test.go b/internal/streamingnode/server/wal/metricsutil/wal_write_trace_test.go new file mode 100644 index 0000000000..4d1681cfd5 --- /dev/null +++ b/internal/streamingnode/server/wal/metricsutil/wal_write_trace_test.go @@ -0,0 +1,92 @@ +//go:build test + +package metricsutil + +import ( + "context" + "encoding/json" + "os" + "path/filepath" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/trace" + + "github.com/milvus-io/milvus-proto/go-api/v3/msgpb" + "github.com/milvus-io/milvus/pkg/v3/mlog" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" + "github.com/milvus-io/milvus/pkg/v3/streaming/util/types" + "github.com/milvus-io/milvus/pkg/v3/streaming/walimpls/impls/walimplstest" + "github.com/milvus-io/milvus/pkg/v3/util/paramtable" +) + +func TestAppendMetricsDoneUsesProvidedTraceContext(t *testing.T) { + paramtable.Init() + + logDir := t.TempDir() + logFile := filepath.Join(logDir, "wal.log") + logger, props, err := mlog.InitLogger(&mlog.Config{ + Level: "debug", + Format: "json", + DisableTimestamp: true, + DisableCaller: true, + DisableStacktrace: true, + File: mlog.FileLogConfig{ + RootPath: logDir, + Filename: "wal.log", + }, + }) + require.NoError(t, err) + mlog.ReplaceGlobals(logger, props) + defer func() { + _ = logger.Sync() + restoreLogger, restoreProps, err := mlog.InitTestLogger(t, &mlog.Config{ + Level: "info", + DisableTimestamp: true, + DisableCaller: true, + DisableStacktrace: true, + }) + require.NoError(t, err) + mlog.ReplaceGlobals(restoreLogger, restoreProps) + }() + + traceID, err := trace.TraceIDFromHex("0102030405060708090a0b0c0d0e0f10") + require.NoError(t, err) + spanID, err := trace.SpanIDFromHex("0102030405060708") + require.NoError(t, err) + ctx := trace.ContextWithSpanContext(context.Background(), trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: traceID, + SpanID: spanID, + })) + + writeMetrics := NewWriteMetrics(types.PChannelInfo{Name: "pchannel-test", Term: 1}, message.WALNameTest) + msg := message.NewTimeTickMessageBuilderV1(). + WithHeader(&message.TimeTickMessageHeader{}). + WithBody(&msgpb.TimeTickMsg{}). + WithAllVChannel(). + MustBuildMutable() + appendMetrics := writeMetrics.StartAppend(msg) + appendMetrics.Done(ctx, &types.AppendResult{ + MessageID: walimplstest.NewTestMessageID(1), + LastConfirmedMessageID: walimplstest.NewTestMessageID(1), + TimeTick: 1, + }, nil) + require.NoError(t, logger.Sync()) + + content, err := os.ReadFile(logFile) + require.NoError(t, err) + var entry map[string]any + for _, line := range strings.Split(strings.TrimSpace(string(content)), "\n") { + var current map[string]any + require.NoError(t, json.Unmarshal([]byte(line), ¤t)) + if current["traceID"] == "0102030405060708090a0b0c0d0e0f10" { + entry = current + break + } + } + require.NotNil(t, entry, string(content)) + assert.Equal(t, "0102030405060708090a0b0c0d0e0f10", entry["traceID"]) + assert.Equal(t, "0102030405060708", entry["spanID"]) +} diff --git a/internal/streamingnode/server/wal/recovery/recovery_drop_collection_test.go b/internal/streamingnode/server/wal/recovery/recovery_drop_collection_test.go index c593453be5..9486cc9af5 100644 --- a/internal/streamingnode/server/wal/recovery/recovery_drop_collection_test.go +++ b/internal/streamingnode/server/wal/recovery/recovery_drop_collection_test.go @@ -3,6 +3,7 @@ package recovery import ( + "context" "testing" "github.com/stretchr/testify/assert" @@ -111,7 +112,7 @@ func TestHandleDropCollection_VChannelAlreadyDropped_FlushesOrphanedSegments(t * // Replay DropCollection for the already-dropped vchannel. dropMsg := buildDropCollectionMsg("v1", 100, 50, 50) - rs.handleDropCollection(dropMsg) + rs.handleDropCollection(context.Background(), dropMsg) // Verify: orphaned segments for collection 100 are flushed. assert.False(t, rs.segments[1001].IsGrowing(), "segment 1001 should be flushed") @@ -129,7 +130,7 @@ func TestHandleDropCollection_VChannelNotFound_FlushesOrphanedSegments(t *testin addGrowingSegment(rs, 1001, 100, 200, "v1") dropMsg := buildDropCollectionMsg("v1", 100, 50, 50) - rs.handleDropCollection(dropMsg) + rs.handleDropCollection(context.Background(), dropMsg) // Segment should be flushed even though vchannel doesn't exist. assert.False(t, rs.segments[1001].IsGrowing(), "segment 1001 should be flushed") @@ -143,7 +144,7 @@ func TestHandleDropCollection_NormalCase_StillWorks(t *testing.T) { addGrowingSegment(rs, 1001, 100, 200, "v1") dropMsg := buildDropCollectionMsg("v1", 100, 50, 50) - rs.handleDropCollection(dropMsg) + rs.handleDropCollection(context.Background(), dropMsg) // vchannel should be marked as DROPPED. assert.Equal(t, streamingpb.VChannelState_VCHANNEL_STATE_DROPPED, rs.vchannels["v1"].meta.State) @@ -225,7 +226,7 @@ func TestHandleCreateSegment_SkipsForDroppedVChannel(t *testing.T) { WithLastConfirmed(rmq.NewRmqID(50)). IntoImmutableMessage(rmq.NewRmqID(50)) - rs.handleCreateSegment(message.MustAsImmutableCreateSegmentMessageV2(createMsg)) + rs.handleCreateSegment(context.Background(), message.MustAsImmutableCreateSegmentMessageV2(createMsg)) // Segment should NOT have been created. assert.Empty(t, rs.segments, "no segment should be created for a dropped vchannel") @@ -250,7 +251,7 @@ func TestHandleCreateSegment_SkipsForNonExistentVChannel(t *testing.T) { WithLastConfirmed(rmq.NewRmqID(50)). IntoImmutableMessage(rmq.NewRmqID(50)) - rs.handleCreateSegment(message.MustAsImmutableCreateSegmentMessageV2(createMsg)) + rs.handleCreateSegment(context.Background(), message.MustAsImmutableCreateSegmentMessageV2(createMsg)) // Segment should NOT have been created. assert.Empty(t, rs.segments, "no segment should be created for a non-existent vchannel") @@ -277,7 +278,7 @@ func TestHandleCreateSegment_NormalCase_StillWorks(t *testing.T) { WithLastConfirmed(rmq.NewRmqID(50)). IntoImmutableMessage(rmq.NewRmqID(50)) - rs.handleCreateSegment(message.MustAsImmutableCreateSegmentMessageV2(createMsg)) + rs.handleCreateSegment(context.Background(), message.MustAsImmutableCreateSegmentMessageV2(createMsg)) // Segment should be created normally. assert.Len(t, rs.segments, 1) @@ -303,7 +304,7 @@ func TestFullReplayScenario_DroppedCollectionReplay(t *testing.T) { WithTimeTick(10). WithLastConfirmed(rmq.NewRmqID(10)). IntoImmutableMessage(rmq.NewRmqID(10)) - rs.handleCreateCollection(message.MustAsImmutableCreateCollectionMessageV1(createCollMsg)) + rs.handleCreateCollection(context.Background(), message.MustAsImmutableCreateCollectionMessageV1(createCollMsg)) // Step 2: CreateSegment replayed createSegMsg := message.NewCreateSegmentMessageBuilderV2(). @@ -320,11 +321,11 @@ func TestFullReplayScenario_DroppedCollectionReplay(t *testing.T) { WithTimeTick(20). WithLastConfirmed(rmq.NewRmqID(20)). IntoImmutableMessage(rmq.NewRmqID(20)) - rs.handleCreateSegment(message.MustAsImmutableCreateSegmentMessageV2(createSegMsg)) + rs.handleCreateSegment(context.Background(), message.MustAsImmutableCreateSegmentMessageV2(createSegMsg)) // Step 3: DropCollection replayed — should flush the segment and mark vchannel dropped. dropMsg := buildDropCollectionMsg("v1", 100, 30, 30) - rs.handleDropCollection(dropMsg) + rs.handleDropCollection(context.Background(), dropMsg) // After drop: vchannel is DROPPED, segment is FLUSHED. assert.Equal(t, streamingpb.VChannelState_VCHANNEL_STATE_DROPPED, rs.vchannels["v1"].meta.State) @@ -350,7 +351,7 @@ func TestFullReplayScenario_PartialEtcdPersist(t *testing.T) { // Then DropCollection is replayed again. dropMsg := buildDropCollectionMsg("v1", 100, 50, 50) - rs.handleDropCollection(dropMsg) + rs.handleDropCollection(context.Background(), dropMsg) // All segments should be flushed. assert.False(t, rs.segments[1001].IsGrowing()) diff --git a/internal/streamingnode/server/wal/recovery/recovery_storage_flushall_test.go b/internal/streamingnode/server/wal/recovery/recovery_storage_flushall_test.go index c083bbff19..c986edb89d 100644 --- a/internal/streamingnode/server/wal/recovery/recovery_storage_flushall_test.go +++ b/internal/streamingnode/server/wal/recovery/recovery_storage_flushall_test.go @@ -1,6 +1,7 @@ package recovery import ( + "context" "testing" "github.com/stretchr/testify/assert" @@ -85,7 +86,7 @@ func TestFlushAllOnControlChannel(t *testing.T) { // Create FlushAll message on control channel via broadcast split flushAllMsg := buildFlushAllOnControlChannel(pchannel, controlChannel, 1, 10, 5, 10) - rs.observeMessage(flushAllMsg) + rs.observeMessage(context.Background(), flushAllMsg) // All segments should be FLUSHED for segID, seg := range rs.segments { @@ -112,7 +113,7 @@ func TestFlushAllOnControlChannel(t *testing.T) { WithLastConfirmed(rmq.NewRmqID(5)). IntoImmutableMessage(rmq.NewRmqID(10)) - rs.observeMessage(flushAllMsg) + rs.observeMessage(context.Background(), flushAllMsg) // All segments should be FLUSHED for segID, seg := range rs.segments { @@ -131,7 +132,7 @@ func TestFlushAllOnControlChannel(t *testing.T) { // First FlushAll flushAllMsg1 := buildFlushAllOnControlChannel(pchannel, controlChannel, 1, 10, 5, 10) - rs.observeMessage(flushAllMsg1) + rs.observeMessage(context.Background(), flushAllMsg1) for _, seg := range rs.segments { assert.Equal(t, streamingpb.SegmentAssignmentState_SEGMENT_ASSIGNMENT_STATE_FLUSHED, seg.meta.State) @@ -140,7 +141,7 @@ func TestFlushAllOnControlChannel(t *testing.T) { // Second FlushAll with higher timetick - should be idempotent flushAllMsg2 := buildFlushAllOnControlChannel(pchannel, controlChannel, 2, 20, 15, 20) - rs.observeMessage(flushAllMsg2) + rs.observeMessage(context.Background(), flushAllMsg2) // Segments should remain FLUSHED for segID, seg := range rs.segments { diff --git a/internal/streamingnode/server/wal/recovery/recovery_storage_impl.go b/internal/streamingnode/server/wal/recovery/recovery_storage_impl.go index bd60185cc0..9be4e03338 100644 --- a/internal/streamingnode/server/wal/recovery/recovery_storage_impl.go +++ b/internal/streamingnode/server/wal/recovery/recovery_storage_impl.go @@ -149,17 +149,19 @@ func (r *recoveryStorageImpl) GetSchema(ctx context.Context, vchannel string, ti } // ObserveMessage is called when a new message is observed. -func (r *recoveryStorageImpl) ObserveMessage(ctx context.Context, msg message.ImmutableMessage) error { +func (r *recoveryStorageImpl) ObserveMessage(ctx context.Context, msg message.ImmutableMessage) (err error) { + ctx = message.ExtractTraceContext(ctx, msg) + if h := msg.BroadcastHeader(); h != nil { if err := streaming.WAL().Broadcast().Ack(ctx, msg); err != nil { - r.Logger().Warn(context.TODO(), "failed to ack broadcast message", mlog.Err(err)) + r.Logger().Warn(ctx, "failed to ack broadcast message", mlog.Err(err)) return err } } r.mu.Lock() defer r.mu.Unlock() - r.observeMessage(msg) + r.observeMessage(ctx, msg) return nil } @@ -222,10 +224,10 @@ func (r *recoveryStorageImpl) consumeDirtySnapshot() *RecoverySnapshot { } // observeMessage observes a message and update the recovery storage. -func (r *recoveryStorageImpl) observeMessage(msg message.ImmutableMessage) { +func (r *recoveryStorageImpl) observeMessage(ctx context.Context, msg message.ImmutableMessage) { if msg.TimeTick() <= r.checkpoint.TimeTick { if r.Logger().Level().Enabled(mlog.DebugLevel) { - r.Logger().Debug(context.TODO(), "skip the message before the checkpoint", + r.Logger().Debug(ctx, "skip the message before the checkpoint", mlog.FieldMessage(msg), mlog.Uint64("checkpoint", r.checkpoint.TimeTick), mlog.Uint64("incoming", msg.TimeTick()), @@ -233,9 +235,9 @@ func (r *recoveryStorageImpl) observeMessage(msg message.ImmutableMessage) { } return } - r.handleMessage(msg) + r.handleMessage(ctx, msg) - r.updateCheckpoint(msg) + r.updateCheckpoint(ctx, msg) r.metrics.ObServeInMemMetrics(r.checkpoint.TimeTick) if !msg.IsPersisted() { @@ -249,7 +251,7 @@ func (r *recoveryStorageImpl) observeMessage(msg message.ImmutableMessage) { } // updateCheckpoint updates the checkpoint of the recovery storage. -func (r *recoveryStorageImpl) updateCheckpoint(msg message.ImmutableMessage) { +func (r *recoveryStorageImpl) updateCheckpoint(ctx context.Context, msg message.ImmutableMessage) { if msg.MessageType() == message.MessageTypeAlterReplicateConfig { cfg := message.MustAsImmutableAlterReplicateConfigMessageV2(msg) header := cfg.Header() @@ -257,7 +259,7 @@ func (r *recoveryStorageImpl) updateCheckpoint(msg message.ImmutableMessage) { // Check ignore field - if true, skip updating ReplicateConfig and ReplicateCheckpoint // This is used for incomplete switchover messages that should be ignored after force promote if header.Ignore { - r.Logger().Info(context.TODO(), "AlterReplicateConfig message has ignore flag set, skipping checkpoint update", + r.Logger().Info(ctx, "AlterReplicateConfig message has ignore flag set, skipping checkpoint update", mlog.Bool("forcePromote", header.ForcePromote)) } else { r.checkpoint.ReplicateConfig = header.ReplicateConfiguration @@ -302,11 +304,11 @@ func (r *recoveryStorageImpl) updateCheckpoint(msg message.ImmutableMessage) { return } if r.checkpoint.ReplicateCheckpoint == nil { - r.detectInconsistency(msg, "replicate checkpoint is nil when incoming replicate message") + r.detectInconsistency(ctx, msg, "replicate checkpoint is nil when incoming replicate message") return } if replicateHeader.ClusterID != r.checkpoint.ReplicateCheckpoint.ClusterID { - r.detectInconsistency(msg, + r.detectInconsistency(ctx, msg, "replicate header cluster id mismatch", mlog.String("expected", r.checkpoint.ReplicateCheckpoint.ClusterID), mlog.String("actual", replicateHeader.ClusterID)) @@ -317,7 +319,7 @@ func (r *recoveryStorageImpl) updateCheckpoint(msg message.ImmutableMessage) { } // The incoming message id is always sorted with timetick. -func (r *recoveryStorageImpl) handleMessage(msg message.ImmutableMessage) { +func (r *recoveryStorageImpl) handleMessage(ctx context.Context, msg message.ImmutableMessage) { if funcutil.IsControlChannel(msg.VChannel()) && !msg.IsPChannelLevel() { // message on control channel except pchannel-level messages is just used to determine the DDL/DCL order, // will not affect the recovery storage, so skip it. @@ -326,66 +328,66 @@ func (r *recoveryStorageImpl) handleMessage(msg message.ImmutableMessage) { if msg.VChannel() != "" && !msg.IsPChannelLevel() && msg.MessageType() != message.MessageTypeCreateCollection && msg.MessageType() != message.MessageTypeDropCollection && r.vchannels[msg.VChannel()] == nil && !funcutil.IsControlChannel(msg.VChannel()) { - r.detectInconsistency(msg, "vchannel not found") + r.detectInconsistency(ctx, msg, "vchannel not found") } switch msg.MessageType() { case message.MessageTypeInsert: immutableMsg := message.MustAsImmutableInsertMessageV1(msg) - r.handleInsert(immutableMsg) + r.handleInsert(ctx, immutableMsg) case message.MessageTypeDelete: immutableMsg := message.MustAsImmutableDeleteMessageV1(msg) r.handleDelete(immutableMsg) case message.MessageTypeCreateSegment: immutableMsg := message.MustAsImmutableCreateSegmentMessageV2(msg) - r.handleCreateSegment(immutableMsg) + r.handleCreateSegment(ctx, immutableMsg) case message.MessageTypeFlush: immutableMsg := message.MustAsImmutableFlushMessageV2(msg) - r.handleFlush(immutableMsg) + r.handleFlush(ctx, immutableMsg) case message.MessageTypeManualFlush: immutableMsg := message.MustAsImmutableManualFlushMessageV2(msg) - r.handleManualFlush(immutableMsg) + r.handleManualFlush(ctx, immutableMsg) case message.MessageTypeFlushAll: immutableMsg := message.MustAsImmutableFlushAllMessageV2(msg) - r.handleFlushAll(immutableMsg) + r.handleFlushAll(ctx, immutableMsg) case message.MessageTypeCreateCollection: immutableMsg := message.MustAsImmutableCreateCollectionMessageV1(msg) - r.handleCreateCollection(immutableMsg) + r.handleCreateCollection(ctx, immutableMsg) case message.MessageTypeDropCollection: immutableMsg := message.MustAsImmutableDropCollectionMessageV1(msg) - r.handleDropCollection(immutableMsg) + r.handleDropCollection(ctx, immutableMsg) case message.MessageTypeCreatePartition: immutableMsg := message.MustAsImmutableCreatePartitionMessageV1(msg) - r.handleCreatePartition(immutableMsg) + r.handleCreatePartition(ctx, immutableMsg) case message.MessageTypeDropPartition: immutableMsg := message.MustAsImmutableDropPartitionMessageV1(msg) - r.handleDropPartition(immutableMsg) + r.handleDropPartition(ctx, immutableMsg) case message.MessageTypeTxn: immutableMsg := message.AsImmutableTxnMessage(msg) - r.handleTxn(immutableMsg) + r.handleTxn(ctx, immutableMsg) case message.MessageTypeImport: immutableMsg := message.MustAsImmutableImportMessageV1(msg) r.handleImport(immutableMsg) case message.MessageTypeSchemaChange: immutableMsg := message.MustAsImmutableSchemaChangeMessageV2(msg) - r.handleSchemaChange(immutableMsg) + r.handleSchemaChange(ctx, immutableMsg) case message.MessageTypeAlterCollection: immutableMsg := message.MustAsImmutableAlterCollectionMessageV2(msg) - r.handleAlterCollection(immutableMsg) + r.handleAlterCollection(ctx, immutableMsg) case message.MessageTypeTruncateCollection: immutableMsg := message.MustAsImmutableTruncateCollectionMessageV2(msg) - r.handleTruncateCollection(immutableMsg) + r.handleTruncateCollection(ctx, immutableMsg) case message.MessageTypeTimeTick: // nothing, the time tick message make no recovery operation. case message.MessageTypeAlterWAL: immutableMsg := message.MustAsImmutableAlterWALMessageV2(msg) - r.handleAlterWAL(immutableMsg) + r.handleAlterWAL(ctx, immutableMsg) } } // handleAlterWAL handles the alter WAL message. // Flushes all growing segments to ensure segment data does not span across different WAL implementations. -func (r *recoveryStorageImpl) handleAlterWAL(msg message.ImmutableAlterWALMessageV2) { +func (r *recoveryStorageImpl) handleAlterWAL(ctx context.Context, msg message.ImmutableAlterWALMessageV2) { header := msg.Header() segmentIDs := make([]int64, 0) @@ -403,14 +405,14 @@ func (r *recoveryStorageImpl) handleAlterWAL(msg message.ImmutableAlterWALMessag } if len(segmentIDs) > 0 { - r.Logger().Info(context.TODO(), "flush all growing segments for WAL switch", + r.Logger().Info(ctx, "flush all growing segments for WAL switch", mlog.FieldMessage(msg), mlog.Stringer("targetWALName", header.TargetWalName), mlog.Int64s("segmentIDs", segmentIDs), mlog.Uint64s("rows", rows), mlog.Uint64s("binarySize", binarySize)) } else { - r.Logger().Info(context.TODO(), "no growing segments to flush for WAL switch", + r.Logger().Info(ctx, "no growing segments to flush for WAL switch", mlog.FieldMessage(msg), mlog.Stringer("targetWALName", header.TargetWalName)) } @@ -425,12 +427,12 @@ func (r *recoveryStorageImpl) handleAlterWAL(msg message.ImmutableAlterWALMessag } // handleInsert handles the insert message. -func (r *recoveryStorageImpl) handleInsert(msg message.ImmutableInsertMessageV1) { +func (r *recoveryStorageImpl) handleInsert(ctx context.Context, msg message.ImmutableInsertMessageV1) { for _, partition := range msg.Header().GetPartitions() { if segment, ok := r.segments[partition.SegmentAssignment.SegmentId]; ok && segment.IsGrowing() { segment.ObserveInsert(msg.TimeTick(), partition) } else { - r.detectInconsistency(msg, "segment not found") + r.detectInconsistency(ctx, msg, "segment not found") } } } @@ -440,12 +442,12 @@ func (r *recoveryStorageImpl) handleDelete(msg message.ImmutableDeleteMessageV1) } // handleCreateSegment handles the create segment message. -func (r *recoveryStorageImpl) handleCreateSegment(msg message.ImmutableCreateSegmentMessageV2) { +func (r *recoveryStorageImpl) handleCreateSegment(ctx context.Context, msg message.ImmutableCreateSegmentMessageV2) { // Skip segment creation if the vchannel does not exist (collection was dropped). // During WAL replay (e.g., Kafka offset reset), CreateSegment messages may appear // for collections whose vchannels have already been cleaned up. if vchannelInfo, ok := r.vchannels[msg.VChannel()]; !ok || vchannelInfo.meta.State == streamingpb.VChannelState_VCHANNEL_STATE_DROPPED { - r.Logger().Warn(context.TODO(), "skip create segment for non-active vchannel", + r.Logger().Warn(ctx, "skip create segment for non-active vchannel", mlog.FieldMessage(msg), mlog.String("vchannel", msg.VChannel()), mlog.Int64("segmentID", msg.Header().SegmentId), @@ -454,37 +456,37 @@ func (r *recoveryStorageImpl) handleCreateSegment(msg message.ImmutableCreateSeg } segment := newSegmentRecoveryInfoFromCreateSegmentMessage(msg) r.segments[segment.meta.SegmentId] = segment - r.Logger().Info(context.TODO(), "create segment", mlog.FieldMessage(msg)) + r.Logger().Info(ctx, "create segment", mlog.FieldMessage(msg)) } // handleFlush handles the flush message. -func (r *recoveryStorageImpl) handleFlush(msg message.ImmutableFlushMessageV2) { +func (r *recoveryStorageImpl) handleFlush(ctx context.Context, msg message.ImmutableFlushMessageV2) { header := msg.Header() if segment, ok := r.segments[header.SegmentId]; ok { segment.ObserveFlush(msg.TimeTick()) - r.Logger().Info(context.TODO(), "flush segment", mlog.FieldMessage(msg), mlog.Uint64("rows", segment.Rows()), mlog.Uint64("binarySize", segment.BinarySize())) + r.Logger().Info(ctx, "flush segment", mlog.FieldMessage(msg), mlog.Uint64("rows", segment.Rows()), mlog.Uint64("binarySize", segment.BinarySize())) } } // handleManualFlush handles the manual flush message. -func (r *recoveryStorageImpl) handleManualFlush(msg message.ImmutableManualFlushMessageV2) { +func (r *recoveryStorageImpl) handleManualFlush(ctx context.Context, msg message.ImmutableManualFlushMessageV2) { segments := make(map[int64]struct{}, len(msg.Header().SegmentIds)) for _, segmentID := range msg.Header().SegmentIds { segments[segmentID] = struct{}{} } - r.flushSegments(msg, segments) + r.flushSegments(ctx, msg, segments) } // handleFlushAll handles the flush all message. -func (r *recoveryStorageImpl) handleFlushAll(msg message.ImmutableFlushAllMessageV2) { +func (r *recoveryStorageImpl) handleFlushAll(ctx context.Context, msg message.ImmutableFlushAllMessageV2) { segments := lo.MapValues(r.segments, func(segment *segmentRecoveryInfo, _ int64) struct{} { return struct{}{} }) - r.flushSegments(msg, segments) + r.flushSegments(ctx, msg, segments) } // flushSegments flushes the segments in the recovery storage. -func (r *recoveryStorageImpl) flushSegments(msg message.ImmutableMessage, sealSegmentIDs map[int64]struct{}) { +func (r *recoveryStorageImpl) flushSegments(ctx context.Context, msg message.ImmutableMessage, sealSegmentIDs map[int64]struct{}) { segmentIDs := make([]int64, 0) rows := make([]uint64, 0) binarySize := make([]uint64, 0) @@ -497,9 +499,9 @@ func (r *recoveryStorageImpl) flushSegments(msg message.ImmutableMessage, sealSe } } if len(segmentIDs) != len(sealSegmentIDs) { - r.detectInconsistency(msg, "flush segments not exist", mlog.Int64s("wanted", lo.Keys(sealSegmentIDs)), mlog.Int64s("actually", segmentIDs)) + r.detectInconsistency(ctx, msg, "flush segments not exist", mlog.Int64s("wanted", lo.Keys(sealSegmentIDs)), mlog.Int64s("actually", segmentIDs)) } - r.Logger().Info(context.TODO(), "flush segments of collection by flush", mlog.FieldMessage(msg), + r.Logger().Info(ctx, "flush segments of collection by flush", mlog.FieldMessage(msg), mlog.Uint64s("rows", rows), mlog.Uint64s("binarySize", binarySize), mlog.Int("flushedSegmentCount", len(segmentIDs)), @@ -507,28 +509,28 @@ func (r *recoveryStorageImpl) flushSegments(msg message.ImmutableMessage, sealSe } // handleCreateCollection handles the create collection message. -func (r *recoveryStorageImpl) handleCreateCollection(msg message.ImmutableCreateCollectionMessageV1) { +func (r *recoveryStorageImpl) handleCreateCollection(ctx context.Context, msg message.ImmutableCreateCollectionMessageV1) { if _, ok := r.vchannels[msg.VChannel()]; ok { return } r.vchannels[msg.VChannel()] = newVChannelRecoveryInfoFromCreateCollectionMessage(msg) - r.Logger().Info(context.TODO(), "create collection", mlog.FieldMessage(msg)) + r.Logger().Info(ctx, "create collection", mlog.FieldMessage(msg)) } // handleDropCollection handles the drop collection message. -func (r *recoveryStorageImpl) handleDropCollection(msg message.ImmutableDropCollectionMessageV1) { +func (r *recoveryStorageImpl) handleDropCollection(ctx context.Context, msg message.ImmutableDropCollectionMessageV1) { // Always flush first: during WAL replay, CreateSegment/Insert messages may have recreated // GROWING segments after the vchannel was marked DROPPED (non-atomic etcd persistence or // Kafka offset compaction). Flushing unconditionally ensures idempotent replay. - r.flushAllSegmentOfCollection(msg, msg.Header().CollectionId) + r.flushAllSegmentOfCollection(ctx, msg, msg.Header().CollectionId) if vchannelInfo, ok := r.vchannels[msg.VChannel()]; ok && vchannelInfo.meta.State != streamingpb.VChannelState_VCHANNEL_STATE_DROPPED { vchannelInfo.ObserveDropCollection(msg) } - r.Logger().Info(context.TODO(), "drop collection", mlog.FieldMessage(msg)) + r.Logger().Info(ctx, "drop collection", mlog.FieldMessage(msg)) } // flushAllSegmentOfCollection flushes all segments of the collection. -func (r *recoveryStorageImpl) flushAllSegmentOfCollection(msg message.ImmutableMessage, collectionID int64) { +func (r *recoveryStorageImpl) flushAllSegmentOfCollection(ctx context.Context, msg message.ImmutableMessage, collectionID int64) { segmentIDs := make([]int64, 0) rows := make([]uint64, 0) for _, segment := range r.segments { @@ -538,31 +540,31 @@ func (r *recoveryStorageImpl) flushAllSegmentOfCollection(msg message.ImmutableM rows = append(rows, segment.Rows()) } } - r.Logger().Info(context.TODO(), "flush all segments of collection", mlog.FieldMessage(msg), mlog.Int64s("segmentIDs", segmentIDs), mlog.Uint64s("rows", rows)) + r.Logger().Info(ctx, "flush all segments of collection", mlog.FieldMessage(msg), mlog.Int64s("segmentIDs", segmentIDs), mlog.Uint64s("rows", rows)) } // handleCreatePartition handles the create partition message. -func (r *recoveryStorageImpl) handleCreatePartition(msg message.ImmutableCreatePartitionMessageV1) { +func (r *recoveryStorageImpl) handleCreatePartition(ctx context.Context, msg message.ImmutableCreatePartitionMessageV1) { if vchannelInfo, ok := r.vchannels[msg.VChannel()]; !ok || vchannelInfo.meta.State == streamingpb.VChannelState_VCHANNEL_STATE_DROPPED { return } r.vchannels[msg.VChannel()].ObserveCreatePartition(msg) - r.Logger().Info(context.TODO(), "create partition", mlog.FieldMessage(msg)) + r.Logger().Info(ctx, "create partition", mlog.FieldMessage(msg)) } // handleDropPartition handles the drop partition message. -func (r *recoveryStorageImpl) handleDropPartition(msg message.ImmutableDropPartitionMessageV1) { +func (r *recoveryStorageImpl) handleDropPartition(ctx context.Context, msg message.ImmutableDropPartitionMessageV1) { // Always flush first: same rationale as handleDropCollection — orphaned GROWING segments // may exist for this partition due to non-atomic etcd persistence or WAL offset reset. - r.flushAllSegmentOfPartition(msg, msg.Header().PartitionId) + r.flushAllSegmentOfPartition(ctx, msg, msg.Header().PartitionId) if vchannelInfo, ok := r.vchannels[msg.VChannel()]; ok && vchannelInfo.meta.State != streamingpb.VChannelState_VCHANNEL_STATE_DROPPED { vchannelInfo.ObserveDropPartition(msg) } - r.Logger().Info(context.TODO(), "drop partition", mlog.FieldMessage(msg)) + r.Logger().Info(ctx, "drop partition", mlog.FieldMessage(msg)) } // flushAllSegmentOfPartition flushes all segments of the partition. -func (r *recoveryStorageImpl) flushAllSegmentOfPartition(msg message.ImmutableMessage, partitionID int64) { +func (r *recoveryStorageImpl) flushAllSegmentOfPartition(ctx context.Context, msg message.ImmutableMessage, partitionID int64) { segmentIDs := make([]int64, 0) rows := make([]uint64, 0) for _, segment := range r.segments { @@ -572,13 +574,13 @@ func (r *recoveryStorageImpl) flushAllSegmentOfPartition(msg message.ImmutableMe rows = append(rows, segment.Rows()) } } - r.Logger().Info(context.TODO(), "flush all segments of partition", mlog.FieldMessage(msg), mlog.Int64s("segmentIDs", segmentIDs), mlog.Uint64s("rows", rows)) + r.Logger().Info(ctx, "flush all segments of partition", mlog.FieldMessage(msg), mlog.Int64s("segmentIDs", segmentIDs), mlog.Uint64s("rows", rows)) } // handleTxn handles the txn message. -func (r *recoveryStorageImpl) handleTxn(msg message.ImmutableTxnMessage) { +func (r *recoveryStorageImpl) handleTxn(ctx context.Context, msg message.ImmutableTxnMessage) { msg.RangeOver(func(im message.ImmutableMessage) error { - r.handleMessage(im) + r.handleMessage(message.ExtractTraceContext(ctx, im), im) return nil }) } @@ -588,13 +590,13 @@ func (r *recoveryStorageImpl) handleImport(_ message.ImmutableImportMessageV1) { } // handleSchemaChange handles the schema change message. -func (r *recoveryStorageImpl) handleSchemaChange(msg message.ImmutableSchemaChangeMessageV2) { +func (r *recoveryStorageImpl) handleSchemaChange(ctx context.Context, msg message.ImmutableSchemaChangeMessageV2) { // when schema change happens, we need to flush all segments in the collection. segments := make(map[int64]struct{}, len(msg.Header().FlushedSegmentIds)) for _, segmentID := range msg.Header().FlushedSegmentIds { segments[segmentID] = struct{}{} } - r.flushSegments(msg, segments) + r.flushSegments(ctx, msg, segments) // persist the schema change into recovery info. if vchannelInfo, ok := r.vchannels[msg.VChannel()]; ok { @@ -603,13 +605,13 @@ func (r *recoveryStorageImpl) handleSchemaChange(msg message.ImmutableSchemaChan } // handlePutCollection handles the put collection message. -func (r *recoveryStorageImpl) handleAlterCollection(msg message.ImmutableAlterCollectionMessageV2) { +func (r *recoveryStorageImpl) handleAlterCollection(ctx context.Context, msg message.ImmutableAlterCollectionMessageV2) { // when put collection happens, we need to flush all segments in the collection. segments := make(map[int64]struct{}, len(msg.Header().FlushedSegmentIds)) for _, segmentID := range msg.Header().FlushedSegmentIds { segments[segmentID] = struct{}{} } - r.flushSegments(msg, segments) + r.flushSegments(ctx, msg, segments) // persist the schema change into recovery info. if vchannelInfo, ok := r.vchannels[msg.VChannel()]; ok { @@ -618,23 +620,23 @@ func (r *recoveryStorageImpl) handleAlterCollection(msg message.ImmutableAlterCo } // handleTruncateCollection handles the truncate collection message. -func (r *recoveryStorageImpl) handleTruncateCollection(msg message.ImmutableTruncateCollectionMessageV2) { +func (r *recoveryStorageImpl) handleTruncateCollection(ctx context.Context, msg message.ImmutableTruncateCollectionMessageV2) { // when truncate collection happens, we need to flush all segments in the collection. segments := make(map[int64]struct{}, len(msg.Header().SegmentIds)) for _, segmentID := range msg.Header().SegmentIds { segments[segmentID] = struct{}{} } - r.flushSegments(msg, segments) + r.flushSegments(ctx, msg, segments) } // detectInconsistency detects the inconsistency in the recovery storage. -func (r *recoveryStorageImpl) detectInconsistency(msg message.ImmutableMessage, reason string, extra ...mlog.Field) { +func (r *recoveryStorageImpl) detectInconsistency(ctx context.Context, msg message.ImmutableMessage, reason string, extra ...mlog.Field) { fields := make([]mlog.Field, 0, len(extra)+2) fields = append(fields, mlog.FieldMessage(msg), mlog.String("reason", reason)) fields = append(fields, extra...) // The log is not fatal in some cases. // because our meta is not atomic-updated, so these error may be logged if crashes when meta updated partially. - r.Logger().Warn(context.TODO(), "inconsistency detected", fields...) + r.Logger().Warn(ctx, "inconsistency detected", fields...) r.metrics.ObserveInconsitentEvent() } diff --git a/internal/streamingnode/server/wal/recovery/recovery_storage_test.go b/internal/streamingnode/server/wal/recovery/recovery_storage_test.go index e9e6dc51c8..a4938cdf46 100644 --- a/internal/streamingnode/server/wal/recovery/recovery_storage_test.go +++ b/internal/streamingnode/server/wal/recovery/recovery_storage_test.go @@ -215,7 +215,7 @@ func TestRecoveryStorageManualFlushMarksSegmentsFlushed(t *testing.T) { WithLastConfirmedUseMessageID(). IntoImmutableMessage(rmq.NewRmqID(2)) - r.handleManualFlush(message.MustAsImmutableManualFlushMessageV2(msg)) + r.handleManualFlush(context.Background(), message.MustAsImmutableManualFlushMessageV2(msg)) segment := r.segments[segmentID] snapshot, shouldBeRemoved := segment.ConsumeDirtyAndGetSnapshot() require.NotNil(t, snapshot) diff --git a/internal/streamingnode/server/wal/recovery/replicate_checkpoint_test.go b/internal/streamingnode/server/wal/recovery/replicate_checkpoint_test.go index b967934691..8f0a7fe838 100644 --- a/internal/streamingnode/server/wal/recovery/replicate_checkpoint_test.go +++ b/internal/streamingnode/server/wal/recovery/replicate_checkpoint_test.go @@ -1,6 +1,7 @@ package recovery import ( + "context" "testing" "github.com/stretchr/testify/assert" @@ -19,11 +20,11 @@ func TestUpdateCheckpoint(t *testing.T) { metrics: newRecoveryStorageMetrics(types.PChannelInfo{Name: "test1-rootcoord-dml_0"}), } - rs.updateCheckpoint(newAlterReplicateConfigMessage("test1", []string{"test2"}, 1, walimplstest.NewTestMessageID(1))) + rs.updateCheckpoint(context.Background(), newAlterReplicateConfigMessage("test1", []string{"test2"}, 1, walimplstest.NewTestMessageID(1))) assert.Nil(t, rs.checkpoint.ReplicateCheckpoint) assert.Equal(t, rs.checkpoint.MessageID, walimplstest.NewTestMessageID(1)) assert.Equal(t, rs.checkpoint.TimeTick, uint64(1)) - rs.updateCheckpoint(newAlterReplicateConfigMessage("test2", []string{"test1"}, 1, walimplstest.NewTestMessageID(1))) + rs.updateCheckpoint(context.Background(), newAlterReplicateConfigMessage("test2", []string{"test1"}, 1, walimplstest.NewTestMessageID(1))) assert.NotNil(t, rs.checkpoint.ReplicateCheckpoint) assert.Equal(t, rs.checkpoint.ReplicateCheckpoint.ClusterID, "test2") assert.Equal(t, rs.checkpoint.ReplicateCheckpoint.PChannel, "test2-rootcoord-dml_0") @@ -42,17 +43,17 @@ func TestUpdateCheckpoint(t *testing.T) { immutableReplicateMsg := replicateMsg.WithTimeTick(4). WithLastConfirmed(walimplstest.NewTestMessageID(11)). IntoImmutableMessage(walimplstest.NewTestMessageID(22)) - rs.updateCheckpoint(immutableReplicateMsg) + rs.updateCheckpoint(context.Background(), immutableReplicateMsg) // update with wrong clusterID. - rs.updateCheckpoint(immutableReplicateMsg) + rs.updateCheckpoint(context.Background(), immutableReplicateMsg) assert.NotNil(t, rs.checkpoint.ReplicateCheckpoint) assert.Equal(t, rs.checkpoint.ReplicateCheckpoint.ClusterID, "test2") assert.Equal(t, rs.checkpoint.ReplicateCheckpoint.PChannel, "test2-rootcoord-dml_0") assert.Nil(t, rs.checkpoint.ReplicateCheckpoint.MessageID) assert.Zero(t, rs.checkpoint.ReplicateCheckpoint.TimeTick) - rs.updateCheckpoint(newAlterReplicateConfigMessage("test3", []string{"test2", "test1"}, 1, walimplstest.NewTestMessageID(1))) + rs.updateCheckpoint(context.Background(), newAlterReplicateConfigMessage("test3", []string{"test2", "test1"}, 1, walimplstest.NewTestMessageID(1))) assert.NotNil(t, rs.checkpoint.ReplicateCheckpoint) assert.Equal(t, rs.checkpoint.ReplicateCheckpoint.ClusterID, "test3") assert.Equal(t, rs.checkpoint.ReplicateCheckpoint.PChannel, "test3-rootcoord-dml_0") @@ -60,7 +61,7 @@ func TestUpdateCheckpoint(t *testing.T) { assert.Zero(t, rs.checkpoint.ReplicateCheckpoint.TimeTick) // update with right clusterID. - rs.updateCheckpoint(immutableReplicateMsg) + rs.updateCheckpoint(context.Background(), immutableReplicateMsg) assert.NotNil(t, rs.checkpoint.ReplicateCheckpoint) assert.Equal(t, rs.checkpoint.MessageID, walimplstest.NewTestMessageID(11)) assert.Equal(t, rs.checkpoint.TimeTick, uint64(4)) @@ -69,9 +70,9 @@ func TestUpdateCheckpoint(t *testing.T) { assert.True(t, rs.checkpoint.ReplicateCheckpoint.MessageID.EQ(walimplstest.NewTestMessageID(10))) assert.Equal(t, rs.checkpoint.ReplicateCheckpoint.TimeTick, uint64(3)) - rs.updateCheckpoint(newAlterReplicateConfigMessage("test1", []string{"test2"}, 1, walimplstest.NewTestMessageID(1))) + rs.updateCheckpoint(context.Background(), newAlterReplicateConfigMessage("test1", []string{"test2"}, 1, walimplstest.NewTestMessageID(1))) assert.Nil(t, rs.checkpoint.ReplicateCheckpoint) - rs.updateCheckpoint(immutableReplicateMsg) + rs.updateCheckpoint(context.Background(), immutableReplicateMsg) } // newAlterReplicateConfigMessage creates a new alter replicate config message. diff --git a/internal/streamingnode/server/wal/recovery/salvage_checkpoint_test.go b/internal/streamingnode/server/wal/recovery/salvage_checkpoint_test.go index 62400c01b6..6cefdee4ad 100644 --- a/internal/streamingnode/server/wal/recovery/salvage_checkpoint_test.go +++ b/internal/streamingnode/server/wal/recovery/salvage_checkpoint_test.go @@ -50,12 +50,12 @@ func TestUpdateCheckpointForcePromote(t *testing.T) { } // Start as secondary of test2 - rs.updateCheckpoint(newAlterReplicateConfigMessage("test2", []string{"test1"}, 2, walimplstest.NewTestMessageID(2))) + rs.updateCheckpoint(context.Background(), newAlterReplicateConfigMessage("test2", []string{"test1"}, 2, walimplstest.NewTestMessageID(2))) assert.NotNil(t, rs.checkpoint.ReplicateCheckpoint) assert.Nil(t, rs.pendingSalvageCheckpoint) // Force promote to primary — should capture the salvage checkpoint - rs.updateCheckpoint(newAlterReplicateConfigMessageWithForcePromote("test1", []string{"test2"}, 3, walimplstest.NewTestMessageID(3))) + rs.updateCheckpoint(context.Background(), newAlterReplicateConfigMessageWithForcePromote("test1", []string{"test2"}, 3, walimplstest.NewTestMessageID(3))) assert.Nil(t, rs.checkpoint.ReplicateCheckpoint) assert.NotNil(t, rs.pendingSalvageCheckpoint) assert.Equal(t, "test2", rs.pendingSalvageCheckpoint.ClusterID) @@ -78,9 +78,9 @@ func TestUpdateCheckpointForcePromote(t *testing.T) { metrics: newRecoveryStorageMetrics(types.PChannelInfo{Name: "test1-rootcoord-dml_0"}), } - rs.updateCheckpoint(newAlterReplicateConfigMessage("test2", []string{"test1"}, 2, walimplstest.NewTestMessageID(2))) + rs.updateCheckpoint(context.Background(), newAlterReplicateConfigMessage("test2", []string{"test1"}, 2, walimplstest.NewTestMessageID(2))) // Normal promote (no ForcePromote flag) - rs.updateCheckpoint(newAlterReplicateConfigMessage("test1", []string{"test2"}, 3, walimplstest.NewTestMessageID(3))) + rs.updateCheckpoint(context.Background(), newAlterReplicateConfigMessage("test1", []string{"test2"}, 3, walimplstest.NewTestMessageID(3))) assert.Nil(t, rs.checkpoint.ReplicateCheckpoint) assert.Nil(t, rs.pendingSalvageCheckpoint) }) @@ -98,7 +98,7 @@ func TestUpdateCheckpointForcePromote(t *testing.T) { metrics: newRecoveryStorageMetrics(types.PChannelInfo{Name: "test1-rootcoord-dml_0"}), } - rs.updateCheckpoint(newAlterReplicateConfigMessageWithForcePromote("test1", []string{"test2"}, 2, walimplstest.NewTestMessageID(2))) + rs.updateCheckpoint(context.Background(), newAlterReplicateConfigMessageWithForcePromote("test1", []string{"test2"}, 2, walimplstest.NewTestMessageID(2))) assert.Nil(t, rs.checkpoint.ReplicateCheckpoint) assert.Nil(t, rs.pendingSalvageCheckpoint) }) diff --git a/pkg/mocks/streaming/util/mock_message/mock_MutableMessage.go b/pkg/mocks/streaming/util/mock_message/mock_MutableMessage.go index a3c2eec800..810ec3f84f 100644 --- a/pkg/mocks/streaming/util/mock_message/mock_MutableMessage.go +++ b/pkg/mocks/streaming/util/mock_message/mock_MutableMessage.go @@ -1236,7 +1236,8 @@ func (_c *MockMutableMessage_WithWALTerm_Call) RunAndReturn(run func(int64) mess func NewMockMutableMessage(t interface { mock.TestingT Cleanup(func()) -}) *MockMutableMessage { +}, +) *MockMutableMessage { mock := &MockMutableMessage{} mock.Mock.Test(t) diff --git a/pkg/proto/messages.proto b/pkg/proto/messages.proto index 557f972502..5fa4a8df7d 100644 --- a/pkg/proto/messages.proto +++ b/pkg/proto/messages.proto @@ -798,6 +798,18 @@ message CipherHeader { int64 payload_bytes = 4; // the size of the payload before encryption } +// TraceContextHeader carries the trace context subset (trace_id, span_id, +// flags) stored on a message. Tracestate is intentionally not persisted. +// Serialized into Properties under reserved key `_tc` so that consumers on +// the other side of an RPC / persistence boundary can stitch spans into the +// correct parent-child tree. +// See docs/agent_guides/streaming-system/wal/tracing.md for details. +message TraceContextHeader { + bytes trace_id = 1; // 16 bytes + bytes span_id = 2; // 8 bytes + uint32 flags = 3; // W3C TraceFlags (sampled bit) +} + // TruncateCollectionMessageHeader is the header of truncate collection message. message TruncateCollectionMessageHeader { int64 db_id = 1; diff --git a/pkg/proto/messagespb/messages.pb.go b/pkg/proto/messagespb/messages.pb.go index 1ca0e74227..6f513d8802 100644 --- a/pkg/proto/messagespb/messages.pb.go +++ b/pkg/proto/messagespb/messages.pb.go @@ -6358,6 +6358,75 @@ func (x *CipherHeader) GetPayloadBytes() int64 { return 0 } +// TraceContextHeader carries the trace context subset (trace_id, span_id, +// flags) stored on a message. Tracestate is intentionally not persisted. +// Serialized into Properties under reserved key `_tc` so that consumers on +// the other side of an RPC / persistence boundary can stitch spans into the +// correct parent-child tree. +// See docs/agent_guides/streaming-system/wal/tracing.md for details. +type TraceContextHeader struct { + state protoimpl.MessageState + sizeCache protoimpl.SizeCache + unknownFields protoimpl.UnknownFields + + TraceId []byte `protobuf:"bytes,1,opt,name=trace_id,json=traceId,proto3" json:"trace_id,omitempty"` // 16 bytes + SpanId []byte `protobuf:"bytes,2,opt,name=span_id,json=spanId,proto3" json:"span_id,omitempty"` // 8 bytes + Flags uint32 `protobuf:"varint,3,opt,name=flags,proto3" json:"flags,omitempty"` // W3C TraceFlags (sampled bit) +} + +func (x *TraceContextHeader) Reset() { + *x = TraceContextHeader{} + if protoimpl.UnsafeEnabled { + mi := &file_messages_proto_msgTypes[113] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) + } +} + +func (x *TraceContextHeader) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*TraceContextHeader) ProtoMessage() {} + +func (x *TraceContextHeader) ProtoReflect() protoreflect.Message { + mi := &file_messages_proto_msgTypes[113] + if protoimpl.UnsafeEnabled && x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use TraceContextHeader.ProtoReflect.Descriptor instead. +func (*TraceContextHeader) Descriptor() ([]byte, []int) { + return file_messages_proto_rawDescGZIP(), []int{113} +} + +func (x *TraceContextHeader) GetTraceId() []byte { + if x != nil { + return x.TraceId + } + return nil +} + +func (x *TraceContextHeader) GetSpanId() []byte { + if x != nil { + return x.SpanId + } + return nil +} + +func (x *TraceContextHeader) GetFlags() uint32 { + if x != nil { + return x.Flags + } + return 0 +} + // TruncateCollectionMessageHeader is the header of truncate collection message. type TruncateCollectionMessageHeader struct { state protoimpl.MessageState @@ -6372,7 +6441,7 @@ type TruncateCollectionMessageHeader struct { func (x *TruncateCollectionMessageHeader) Reset() { *x = TruncateCollectionMessageHeader{} if protoimpl.UnsafeEnabled { - mi := &file_messages_proto_msgTypes[113] + mi := &file_messages_proto_msgTypes[114] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -6385,7 +6454,7 @@ func (x *TruncateCollectionMessageHeader) String() string { func (*TruncateCollectionMessageHeader) ProtoMessage() {} func (x *TruncateCollectionMessageHeader) ProtoReflect() protoreflect.Message { - mi := &file_messages_proto_msgTypes[113] + mi := &file_messages_proto_msgTypes[114] if protoimpl.UnsafeEnabled && x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -6398,7 +6467,7 @@ func (x *TruncateCollectionMessageHeader) ProtoReflect() protoreflect.Message { // Deprecated: Use TruncateCollectionMessageHeader.ProtoReflect.Descriptor instead. func (*TruncateCollectionMessageHeader) Descriptor() ([]byte, []int) { - return file_messages_proto_rawDescGZIP(), []int{113} + return file_messages_proto_rawDescGZIP(), []int{114} } func (x *TruncateCollectionMessageHeader) GetDbId() int64 { @@ -6432,7 +6501,7 @@ type TruncateCollectionMessageBody struct { func (x *TruncateCollectionMessageBody) Reset() { *x = TruncateCollectionMessageBody{} if protoimpl.UnsafeEnabled { - mi := &file_messages_proto_msgTypes[114] + mi := &file_messages_proto_msgTypes[115] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -6445,7 +6514,7 @@ func (x *TruncateCollectionMessageBody) String() string { func (*TruncateCollectionMessageBody) ProtoMessage() {} func (x *TruncateCollectionMessageBody) ProtoReflect() protoreflect.Message { - mi := &file_messages_proto_msgTypes[114] + mi := &file_messages_proto_msgTypes[115] if protoimpl.UnsafeEnabled && x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -6458,7 +6527,7 @@ func (x *TruncateCollectionMessageBody) ProtoReflect() protoreflect.Message { // Deprecated: Use TruncateCollectionMessageBody.ProtoReflect.Descriptor instead. func (*TruncateCollectionMessageBody) Descriptor() ([]byte, []int) { - return file_messages_proto_rawDescGZIP(), []int{114} + return file_messages_proto_rawDescGZIP(), []int{115} } // BatchUpdateManifestMessageHeader is the header of batch update manifest message. @@ -6473,7 +6542,7 @@ type BatchUpdateManifestMessageHeader struct { func (x *BatchUpdateManifestMessageHeader) Reset() { *x = BatchUpdateManifestMessageHeader{} if protoimpl.UnsafeEnabled { - mi := &file_messages_proto_msgTypes[115] + mi := &file_messages_proto_msgTypes[116] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -6486,7 +6555,7 @@ func (x *BatchUpdateManifestMessageHeader) String() string { func (*BatchUpdateManifestMessageHeader) ProtoMessage() {} func (x *BatchUpdateManifestMessageHeader) ProtoReflect() protoreflect.Message { - mi := &file_messages_proto_msgTypes[115] + mi := &file_messages_proto_msgTypes[116] if protoimpl.UnsafeEnabled && x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -6499,7 +6568,7 @@ func (x *BatchUpdateManifestMessageHeader) ProtoReflect() protoreflect.Message { // Deprecated: Use BatchUpdateManifestMessageHeader.ProtoReflect.Descriptor instead. func (*BatchUpdateManifestMessageHeader) Descriptor() ([]byte, []int) { - return file_messages_proto_rawDescGZIP(), []int{115} + return file_messages_proto_rawDescGZIP(), []int{116} } func (x *BatchUpdateManifestMessageHeader) GetCollectionId() int64 { @@ -6521,7 +6590,7 @@ type BatchUpdateManifestMessageBody struct { func (x *BatchUpdateManifestMessageBody) Reset() { *x = BatchUpdateManifestMessageBody{} if protoimpl.UnsafeEnabled { - mi := &file_messages_proto_msgTypes[116] + mi := &file_messages_proto_msgTypes[117] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -6534,7 +6603,7 @@ func (x *BatchUpdateManifestMessageBody) String() string { func (*BatchUpdateManifestMessageBody) ProtoMessage() {} func (x *BatchUpdateManifestMessageBody) ProtoReflect() protoreflect.Message { - mi := &file_messages_proto_msgTypes[116] + mi := &file_messages_proto_msgTypes[117] if protoimpl.UnsafeEnabled && x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -6547,7 +6616,7 @@ func (x *BatchUpdateManifestMessageBody) ProtoReflect() protoreflect.Message { // Deprecated: Use BatchUpdateManifestMessageBody.ProtoReflect.Descriptor instead. func (*BatchUpdateManifestMessageBody) Descriptor() ([]byte, []int) { - return file_messages_proto_rawDescGZIP(), []int{116} + return file_messages_proto_rawDescGZIP(), []int{117} } func (x *BatchUpdateManifestMessageBody) GetItems() []*BatchUpdateManifestItem { @@ -6575,7 +6644,7 @@ type BatchUpdateManifestItem struct { func (x *BatchUpdateManifestItem) Reset() { *x = BatchUpdateManifestItem{} if protoimpl.UnsafeEnabled { - mi := &file_messages_proto_msgTypes[117] + mi := &file_messages_proto_msgTypes[118] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -6588,7 +6657,7 @@ func (x *BatchUpdateManifestItem) String() string { func (*BatchUpdateManifestItem) ProtoMessage() {} func (x *BatchUpdateManifestItem) ProtoReflect() protoreflect.Message { - mi := &file_messages_proto_msgTypes[117] + mi := &file_messages_proto_msgTypes[118] if protoimpl.UnsafeEnabled && x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -6601,7 +6670,7 @@ func (x *BatchUpdateManifestItem) ProtoReflect() protoreflect.Message { // Deprecated: Use BatchUpdateManifestItem.ProtoReflect.Descriptor instead. func (*BatchUpdateManifestItem) Descriptor() ([]byte, []int) { - return file_messages_proto_rawDescGZIP(), []int{117} + return file_messages_proto_rawDescGZIP(), []int{118} } func (x *BatchUpdateManifestItem) GetSegmentId() int64 { @@ -6638,7 +6707,7 @@ type BatchUpdateManifestV2ColumnGroups struct { func (x *BatchUpdateManifestV2ColumnGroups) Reset() { *x = BatchUpdateManifestV2ColumnGroups{} if protoimpl.UnsafeEnabled { - mi := &file_messages_proto_msgTypes[118] + mi := &file_messages_proto_msgTypes[119] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -6651,7 +6720,7 @@ func (x *BatchUpdateManifestV2ColumnGroups) String() string { func (*BatchUpdateManifestV2ColumnGroups) ProtoMessage() {} func (x *BatchUpdateManifestV2ColumnGroups) ProtoReflect() protoreflect.Message { - mi := &file_messages_proto_msgTypes[118] + mi := &file_messages_proto_msgTypes[119] if protoimpl.UnsafeEnabled && x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -6664,7 +6733,7 @@ func (x *BatchUpdateManifestV2ColumnGroups) ProtoReflect() protoreflect.Message // Deprecated: Use BatchUpdateManifestV2ColumnGroups.ProtoReflect.Descriptor instead. func (*BatchUpdateManifestV2ColumnGroups) Descriptor() ([]byte, []int) { - return file_messages_proto_rawDescGZIP(), []int{118} + return file_messages_proto_rawDescGZIP(), []int{119} } func (x *BatchUpdateManifestV2ColumnGroups) GetColumnGroups() map[int64]*datapb.FieldBinlog { @@ -7406,7 +7475,13 @@ var file_messages_proto_rawDesc = []byte{ 0x5f, 0x6b, 0x65, 0x79, 0x18, 0x03, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x07, 0x73, 0x61, 0x66, 0x65, 0x4b, 0x65, 0x79, 0x12, 0x23, 0x0a, 0x0d, 0x70, 0x61, 0x79, 0x6c, 0x6f, 0x61, 0x64, 0x5f, 0x62, 0x79, 0x74, 0x65, 0x73, 0x18, 0x04, 0x20, 0x01, 0x28, 0x03, 0x52, 0x0c, 0x70, 0x61, 0x79, 0x6c, - 0x6f, 0x61, 0x64, 0x42, 0x79, 0x74, 0x65, 0x73, 0x22, 0x7c, 0x0a, 0x1f, 0x54, 0x72, 0x75, 0x6e, + 0x6f, 0x61, 0x64, 0x42, 0x79, 0x74, 0x65, 0x73, 0x22, 0x5e, 0x0a, 0x12, 0x54, 0x72, 0x61, 0x63, + 0x65, 0x43, 0x6f, 0x6e, 0x74, 0x65, 0x78, 0x74, 0x48, 0x65, 0x61, 0x64, 0x65, 0x72, 0x12, 0x19, + 0x0a, 0x08, 0x74, 0x72, 0x61, 0x63, 0x65, 0x5f, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0c, + 0x52, 0x07, 0x74, 0x72, 0x61, 0x63, 0x65, 0x49, 0x64, 0x12, 0x17, 0x0a, 0x07, 0x73, 0x70, 0x61, + 0x6e, 0x5f, 0x69, 0x64, 0x18, 0x02, 0x20, 0x01, 0x28, 0x0c, 0x52, 0x06, 0x73, 0x70, 0x61, 0x6e, + 0x49, 0x64, 0x12, 0x14, 0x0a, 0x05, 0x66, 0x6c, 0x61, 0x67, 0x73, 0x18, 0x03, 0x20, 0x01, 0x28, + 0x0d, 0x52, 0x05, 0x66, 0x6c, 0x61, 0x67, 0x73, 0x22, 0x7c, 0x0a, 0x1f, 0x54, 0x72, 0x75, 0x6e, 0x63, 0x61, 0x74, 0x65, 0x43, 0x6f, 0x6c, 0x6c, 0x65, 0x63, 0x74, 0x69, 0x6f, 0x6e, 0x4d, 0x65, 0x73, 0x73, 0x61, 0x67, 0x65, 0x48, 0x65, 0x61, 0x64, 0x65, 0x72, 0x12, 0x13, 0x0a, 0x05, 0x64, 0x62, 0x5f, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x03, 0x52, 0x04, 0x64, 0x62, 0x49, 0x64, @@ -7559,7 +7634,7 @@ func file_messages_proto_rawDescGZIP() []byte { } var file_messages_proto_enumTypes = make([]protoimpl.EnumInfo, 3) -var file_messages_proto_msgTypes = make([]protoimpl.MessageInfo, 124) +var file_messages_proto_msgTypes = make([]protoimpl.MessageInfo, 125) var file_messages_proto_goTypes = []interface{}{ (MessageType)(0), // 0: milvus.proto.messages.MessageType (TxnState)(0), // 1: milvus.proto.messages.TxnState @@ -7677,91 +7752,92 @@ var file_messages_proto_goTypes = []interface{}{ (*ReplicateHeader)(nil), // 113: milvus.proto.messages.ReplicateHeader (*ResourceKey)(nil), // 114: milvus.proto.messages.ResourceKey (*CipherHeader)(nil), // 115: milvus.proto.messages.CipherHeader - (*TruncateCollectionMessageHeader)(nil), // 116: milvus.proto.messages.TruncateCollectionMessageHeader - (*TruncateCollectionMessageBody)(nil), // 117: milvus.proto.messages.TruncateCollectionMessageBody - (*BatchUpdateManifestMessageHeader)(nil), // 118: milvus.proto.messages.BatchUpdateManifestMessageHeader - (*BatchUpdateManifestMessageBody)(nil), // 119: milvus.proto.messages.BatchUpdateManifestMessageBody - (*BatchUpdateManifestItem)(nil), // 120: milvus.proto.messages.BatchUpdateManifestItem - (*BatchUpdateManifestV2ColumnGroups)(nil), // 121: milvus.proto.messages.BatchUpdateManifestV2ColumnGroups - nil, // 122: milvus.proto.messages.Message.PropertiesEntry - nil, // 123: milvus.proto.messages.AlterResourceGroupMessageHeader.ResourceGroupConfigsEntry - nil, // 124: milvus.proto.messages.AlterWALMessageHeader.ConfigEntry - nil, // 125: milvus.proto.messages.RMQMessageLayout.PropertiesEntry - nil, // 126: milvus.proto.messages.BatchUpdateManifestV2ColumnGroups.ColumnGroupsEntry - (datapb.SegmentLevel)(0), // 127: milvus.proto.data.SegmentLevel - (*commonpb.ReplicateConfiguration)(nil), // 128: milvus.proto.common.ReplicateConfiguration - (*schemapb.CollectionSchema)(nil), // 129: milvus.proto.schema.CollectionSchema - (*fieldmaskpb.FieldMask)(nil), // 130: google.protobuf.FieldMask - (commonpb.ConsistencyLevel)(0), // 131: milvus.proto.common.ConsistencyLevel - (*commonpb.KeyValuePair)(nil), // 132: milvus.proto.common.KeyValuePair - (*indexpb.FieldIndex)(nil), // 133: milvus.proto.index.FieldIndex - (commonpb.LoadPriority)(0), // 134: milvus.proto.common.LoadPriority - (*milvuspb.UserEntity)(nil), // 135: milvus.proto.milvus.UserEntity - (*internalpb.CredentialInfo)(nil), // 136: milvus.proto.internal.CredentialInfo - (*milvuspb.RoleEntity)(nil), // 137: milvus.proto.milvus.RoleEntity - (*milvuspb.RBACMeta)(nil), // 138: milvus.proto.milvus.RBACMeta - (*milvuspb.GrantEntity)(nil), // 139: milvus.proto.milvus.GrantEntity - (*milvuspb.PrivilegeGroupInfo)(nil), // 140: milvus.proto.milvus.PrivilegeGroupInfo - (commonpb.WALName)(0), // 141: milvus.proto.common.WALName - (commonpb.MsgType)(0), // 142: milvus.proto.common.MsgType - (*commonpb.MessageID)(nil), // 143: milvus.proto.common.MessageID - (*rgpb.ResourceGroupConfig)(nil), // 144: milvus.proto.rg.ResourceGroupConfig - (*datapb.FieldBinlog)(nil), // 145: milvus.proto.data.FieldBinlog + (*TraceContextHeader)(nil), // 116: milvus.proto.messages.TraceContextHeader + (*TruncateCollectionMessageHeader)(nil), // 117: milvus.proto.messages.TruncateCollectionMessageHeader + (*TruncateCollectionMessageBody)(nil), // 118: milvus.proto.messages.TruncateCollectionMessageBody + (*BatchUpdateManifestMessageHeader)(nil), // 119: milvus.proto.messages.BatchUpdateManifestMessageHeader + (*BatchUpdateManifestMessageBody)(nil), // 120: milvus.proto.messages.BatchUpdateManifestMessageBody + (*BatchUpdateManifestItem)(nil), // 121: milvus.proto.messages.BatchUpdateManifestItem + (*BatchUpdateManifestV2ColumnGroups)(nil), // 122: milvus.proto.messages.BatchUpdateManifestV2ColumnGroups + nil, // 123: milvus.proto.messages.Message.PropertiesEntry + nil, // 124: milvus.proto.messages.AlterResourceGroupMessageHeader.ResourceGroupConfigsEntry + nil, // 125: milvus.proto.messages.AlterWALMessageHeader.ConfigEntry + nil, // 126: milvus.proto.messages.RMQMessageLayout.PropertiesEntry + nil, // 127: milvus.proto.messages.BatchUpdateManifestV2ColumnGroups.ColumnGroupsEntry + (datapb.SegmentLevel)(0), // 128: milvus.proto.data.SegmentLevel + (*commonpb.ReplicateConfiguration)(nil), // 129: milvus.proto.common.ReplicateConfiguration + (*schemapb.CollectionSchema)(nil), // 130: milvus.proto.schema.CollectionSchema + (*fieldmaskpb.FieldMask)(nil), // 131: google.protobuf.FieldMask + (commonpb.ConsistencyLevel)(0), // 132: milvus.proto.common.ConsistencyLevel + (*commonpb.KeyValuePair)(nil), // 133: milvus.proto.common.KeyValuePair + (*indexpb.FieldIndex)(nil), // 134: milvus.proto.index.FieldIndex + (commonpb.LoadPriority)(0), // 135: milvus.proto.common.LoadPriority + (*milvuspb.UserEntity)(nil), // 136: milvus.proto.milvus.UserEntity + (*internalpb.CredentialInfo)(nil), // 137: milvus.proto.internal.CredentialInfo + (*milvuspb.RoleEntity)(nil), // 138: milvus.proto.milvus.RoleEntity + (*milvuspb.RBACMeta)(nil), // 139: milvus.proto.milvus.RBACMeta + (*milvuspb.GrantEntity)(nil), // 140: milvus.proto.milvus.GrantEntity + (*milvuspb.PrivilegeGroupInfo)(nil), // 141: milvus.proto.milvus.PrivilegeGroupInfo + (commonpb.WALName)(0), // 142: milvus.proto.common.WALName + (commonpb.MsgType)(0), // 143: milvus.proto.common.MsgType + (*commonpb.MessageID)(nil), // 144: milvus.proto.common.MessageID + (*rgpb.ResourceGroupConfig)(nil), // 145: milvus.proto.rg.ResourceGroupConfig + (*datapb.FieldBinlog)(nil), // 146: milvus.proto.data.FieldBinlog } var file_messages_proto_depIdxs = []int32{ - 122, // 0: milvus.proto.messages.Message.properties:type_name -> milvus.proto.messages.Message.PropertiesEntry + 123, // 0: milvus.proto.messages.Message.properties:type_name -> milvus.proto.messages.Message.PropertiesEntry 3, // 1: milvus.proto.messages.TxnMessageBody.messages:type_name -> milvus.proto.messages.Message 13, // 2: milvus.proto.messages.InsertMessageHeader.partitions:type_name -> milvus.proto.messages.PartitionSegmentAssignment 14, // 3: milvus.proto.messages.PartitionSegmentAssignment.segment_assignment:type_name -> milvus.proto.messages.SegmentAssignment - 127, // 4: milvus.proto.messages.CreateSegmentMessageHeader.level:type_name -> milvus.proto.data.SegmentLevel - 128, // 5: milvus.proto.messages.AlterReplicateConfigMessageHeader.replicate_configuration:type_name -> milvus.proto.common.ReplicateConfiguration - 129, // 6: milvus.proto.messages.SchemaChangeMessageBody.schema:type_name -> milvus.proto.schema.CollectionSchema - 130, // 7: milvus.proto.messages.AlterCollectionMessageHeader.update_mask:type_name -> google.protobuf.FieldMask + 128, // 4: milvus.proto.messages.CreateSegmentMessageHeader.level:type_name -> milvus.proto.data.SegmentLevel + 129, // 5: milvus.proto.messages.AlterReplicateConfigMessageHeader.replicate_configuration:type_name -> milvus.proto.common.ReplicateConfiguration + 130, // 6: milvus.proto.messages.SchemaChangeMessageBody.schema:type_name -> milvus.proto.schema.CollectionSchema + 131, // 7: milvus.proto.messages.AlterCollectionMessageHeader.update_mask:type_name -> google.protobuf.FieldMask 104, // 8: milvus.proto.messages.AlterCollectionMessageHeader.cache_expirations:type_name -> milvus.proto.messages.CacheExpirations 34, // 9: milvus.proto.messages.AlterCollectionMessageBody.updates:type_name -> milvus.proto.messages.AlterCollectionMessageUpdates - 129, // 10: milvus.proto.messages.AlterCollectionMessageUpdates.schema:type_name -> milvus.proto.schema.CollectionSchema - 131, // 11: milvus.proto.messages.AlterCollectionMessageUpdates.consistency_level:type_name -> milvus.proto.common.ConsistencyLevel - 132, // 12: milvus.proto.messages.AlterCollectionMessageUpdates.properties:type_name -> milvus.proto.common.KeyValuePair + 130, // 10: milvus.proto.messages.AlterCollectionMessageUpdates.schema:type_name -> milvus.proto.schema.CollectionSchema + 132, // 11: milvus.proto.messages.AlterCollectionMessageUpdates.consistency_level:type_name -> milvus.proto.common.ConsistencyLevel + 133, // 12: milvus.proto.messages.AlterCollectionMessageUpdates.properties:type_name -> milvus.proto.common.KeyValuePair 35, // 13: milvus.proto.messages.AlterCollectionMessageUpdates.alter_load_config:type_name -> milvus.proto.messages.AlterLoadConfigOfAlterCollection - 133, // 14: milvus.proto.messages.AlterCollectionMessageUpdates.bound_field_indexes:type_name -> milvus.proto.index.FieldIndex + 134, // 14: milvus.proto.messages.AlterCollectionMessageUpdates.bound_field_indexes:type_name -> milvus.proto.index.FieldIndex 38, // 15: milvus.proto.messages.AlterLoadConfigMessageHeader.load_fields:type_name -> milvus.proto.messages.LoadFieldConfig 39, // 16: milvus.proto.messages.AlterLoadConfigMessageHeader.replicas:type_name -> milvus.proto.messages.LoadReplicaConfig - 134, // 17: milvus.proto.messages.LoadReplicaConfig.priority:type_name -> milvus.proto.common.LoadPriority - 132, // 18: milvus.proto.messages.CreateDatabaseMessageBody.properties:type_name -> milvus.proto.common.KeyValuePair - 132, // 19: milvus.proto.messages.AlterDatabaseMessageBody.properties:type_name -> milvus.proto.common.KeyValuePair + 135, // 17: milvus.proto.messages.LoadReplicaConfig.priority:type_name -> milvus.proto.common.LoadPriority + 133, // 18: milvus.proto.messages.CreateDatabaseMessageBody.properties:type_name -> milvus.proto.common.KeyValuePair + 133, // 19: milvus.proto.messages.AlterDatabaseMessageBody.properties:type_name -> milvus.proto.common.KeyValuePair 46, // 20: milvus.proto.messages.AlterDatabaseMessageBody.alter_load_config:type_name -> milvus.proto.messages.AlterLoadConfigOfAlterDatabase - 135, // 21: milvus.proto.messages.CreateUserMessageHeader.user_entity:type_name -> milvus.proto.milvus.UserEntity - 136, // 22: milvus.proto.messages.CreateUserMessageBody.credential_info:type_name -> milvus.proto.internal.CredentialInfo - 135, // 23: milvus.proto.messages.AlterUserMessageHeader.user_entity:type_name -> milvus.proto.milvus.UserEntity - 136, // 24: milvus.proto.messages.AlterUserMessageBody.credential_info:type_name -> milvus.proto.internal.CredentialInfo - 137, // 25: milvus.proto.messages.AlterRoleMessageHeader.role_entity:type_name -> milvus.proto.milvus.RoleEntity - 135, // 26: milvus.proto.messages.RoleBinding.user_entity:type_name -> milvus.proto.milvus.UserEntity - 137, // 27: milvus.proto.messages.RoleBinding.role_entity:type_name -> milvus.proto.milvus.RoleEntity + 136, // 21: milvus.proto.messages.CreateUserMessageHeader.user_entity:type_name -> milvus.proto.milvus.UserEntity + 137, // 22: milvus.proto.messages.CreateUserMessageBody.credential_info:type_name -> milvus.proto.internal.CredentialInfo + 136, // 23: milvus.proto.messages.AlterUserMessageHeader.user_entity:type_name -> milvus.proto.milvus.UserEntity + 137, // 24: milvus.proto.messages.AlterUserMessageBody.credential_info:type_name -> milvus.proto.internal.CredentialInfo + 138, // 25: milvus.proto.messages.AlterRoleMessageHeader.role_entity:type_name -> milvus.proto.milvus.RoleEntity + 136, // 26: milvus.proto.messages.RoleBinding.user_entity:type_name -> milvus.proto.milvus.UserEntity + 138, // 27: milvus.proto.messages.RoleBinding.role_entity:type_name -> milvus.proto.milvus.RoleEntity 63, // 28: milvus.proto.messages.AlterUserRoleMessageHeader.role_binding:type_name -> milvus.proto.messages.RoleBinding 63, // 29: milvus.proto.messages.DropUserRoleMessageHeader.role_binding:type_name -> milvus.proto.messages.RoleBinding - 138, // 30: milvus.proto.messages.RestoreRBACMessageBody.rbac_meta:type_name -> milvus.proto.milvus.RBACMeta - 139, // 31: milvus.proto.messages.AlterPrivilegeMessageHeader.entity:type_name -> milvus.proto.milvus.GrantEntity - 139, // 32: milvus.proto.messages.DropPrivilegeMessageHeader.entity:type_name -> milvus.proto.milvus.GrantEntity - 140, // 33: milvus.proto.messages.AlterPrivilegeGroupMessageHeader.privilege_group_info:type_name -> milvus.proto.milvus.PrivilegeGroupInfo - 140, // 34: milvus.proto.messages.DropPrivilegeGroupMessageHeader.privilege_group_info:type_name -> milvus.proto.milvus.PrivilegeGroupInfo - 123, // 35: milvus.proto.messages.AlterResourceGroupMessageHeader.resource_group_configs:type_name -> milvus.proto.messages.AlterResourceGroupMessageHeader.ResourceGroupConfigsEntry - 133, // 36: milvus.proto.messages.CreateIndexMessageBody.field_index:type_name -> milvus.proto.index.FieldIndex - 133, // 37: milvus.proto.messages.AlterIndexMessageBody.field_indexes:type_name -> milvus.proto.index.FieldIndex - 141, // 38: milvus.proto.messages.AlterWALMessageHeader.target_wal_name:type_name -> milvus.proto.common.WALName - 124, // 39: milvus.proto.messages.AlterWALMessageHeader.config:type_name -> milvus.proto.messages.AlterWALMessageHeader.ConfigEntry + 139, // 30: milvus.proto.messages.RestoreRBACMessageBody.rbac_meta:type_name -> milvus.proto.milvus.RBACMeta + 140, // 31: milvus.proto.messages.AlterPrivilegeMessageHeader.entity:type_name -> milvus.proto.milvus.GrantEntity + 140, // 32: milvus.proto.messages.DropPrivilegeMessageHeader.entity:type_name -> milvus.proto.milvus.GrantEntity + 141, // 33: milvus.proto.messages.AlterPrivilegeGroupMessageHeader.privilege_group_info:type_name -> milvus.proto.milvus.PrivilegeGroupInfo + 141, // 34: milvus.proto.messages.DropPrivilegeGroupMessageHeader.privilege_group_info:type_name -> milvus.proto.milvus.PrivilegeGroupInfo + 124, // 35: milvus.proto.messages.AlterResourceGroupMessageHeader.resource_group_configs:type_name -> milvus.proto.messages.AlterResourceGroupMessageHeader.ResourceGroupConfigsEntry + 134, // 36: milvus.proto.messages.CreateIndexMessageBody.field_index:type_name -> milvus.proto.index.FieldIndex + 134, // 37: milvus.proto.messages.AlterIndexMessageBody.field_indexes:type_name -> milvus.proto.index.FieldIndex + 142, // 38: milvus.proto.messages.AlterWALMessageHeader.target_wal_name:type_name -> milvus.proto.common.WALName + 125, // 39: milvus.proto.messages.AlterWALMessageHeader.config:type_name -> milvus.proto.messages.AlterWALMessageHeader.ConfigEntry 105, // 40: milvus.proto.messages.CacheExpirations.cache_expirations:type_name -> milvus.proto.messages.CacheExpiration 106, // 41: milvus.proto.messages.CacheExpiration.legacy_proxy_collection_meta_cache:type_name -> milvus.proto.messages.LegacyProxyCollectionMetaCache - 142, // 42: milvus.proto.messages.LegacyProxyCollectionMetaCache.msg_type:type_name -> milvus.proto.common.MsgType - 125, // 43: milvus.proto.messages.RMQMessageLayout.properties:type_name -> milvus.proto.messages.RMQMessageLayout.PropertiesEntry + 143, // 42: milvus.proto.messages.LegacyProxyCollectionMetaCache.msg_type:type_name -> milvus.proto.common.MsgType + 126, // 43: milvus.proto.messages.RMQMessageLayout.properties:type_name -> milvus.proto.messages.RMQMessageLayout.PropertiesEntry 114, // 44: milvus.proto.messages.BroadcastHeader.Resource_keys:type_name -> milvus.proto.messages.ResourceKey - 143, // 45: milvus.proto.messages.ReplicateHeader.message_id:type_name -> milvus.proto.common.MessageID - 143, // 46: milvus.proto.messages.ReplicateHeader.last_confirmed_message_id:type_name -> milvus.proto.common.MessageID + 144, // 45: milvus.proto.messages.ReplicateHeader.message_id:type_name -> milvus.proto.common.MessageID + 144, // 46: milvus.proto.messages.ReplicateHeader.last_confirmed_message_id:type_name -> milvus.proto.common.MessageID 2, // 47: milvus.proto.messages.ResourceKey.domain:type_name -> milvus.proto.messages.ResourceDomain - 120, // 48: milvus.proto.messages.BatchUpdateManifestMessageBody.items:type_name -> milvus.proto.messages.BatchUpdateManifestItem - 121, // 49: milvus.proto.messages.BatchUpdateManifestItem.v2_column_groups:type_name -> milvus.proto.messages.BatchUpdateManifestV2ColumnGroups - 126, // 50: milvus.proto.messages.BatchUpdateManifestV2ColumnGroups.column_groups:type_name -> milvus.proto.messages.BatchUpdateManifestV2ColumnGroups.ColumnGroupsEntry - 144, // 51: milvus.proto.messages.AlterResourceGroupMessageHeader.ResourceGroupConfigsEntry.value:type_name -> milvus.proto.rg.ResourceGroupConfig - 145, // 52: milvus.proto.messages.BatchUpdateManifestV2ColumnGroups.ColumnGroupsEntry.value:type_name -> milvus.proto.data.FieldBinlog + 121, // 48: milvus.proto.messages.BatchUpdateManifestMessageBody.items:type_name -> milvus.proto.messages.BatchUpdateManifestItem + 122, // 49: milvus.proto.messages.BatchUpdateManifestItem.v2_column_groups:type_name -> milvus.proto.messages.BatchUpdateManifestV2ColumnGroups + 127, // 50: milvus.proto.messages.BatchUpdateManifestV2ColumnGroups.column_groups:type_name -> milvus.proto.messages.BatchUpdateManifestV2ColumnGroups.ColumnGroupsEntry + 145, // 51: milvus.proto.messages.AlterResourceGroupMessageHeader.ResourceGroupConfigsEntry.value:type_name -> milvus.proto.rg.ResourceGroupConfig + 146, // 52: milvus.proto.messages.BatchUpdateManifestV2ColumnGroups.ColumnGroupsEntry.value:type_name -> milvus.proto.data.FieldBinlog 53, // [53:53] is the sub-list for method output_type 53, // [53:53] is the sub-list for method input_type 53, // [53:53] is the sub-list for extension type_name @@ -9132,7 +9208,7 @@ func file_messages_proto_init() { } } file_messages_proto_msgTypes[113].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*TruncateCollectionMessageHeader); i { + switch v := v.(*TraceContextHeader); i { case 0: return &v.state case 1: @@ -9144,7 +9220,7 @@ func file_messages_proto_init() { } } file_messages_proto_msgTypes[114].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*TruncateCollectionMessageBody); i { + switch v := v.(*TruncateCollectionMessageHeader); i { case 0: return &v.state case 1: @@ -9156,7 +9232,7 @@ func file_messages_proto_init() { } } file_messages_proto_msgTypes[115].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*BatchUpdateManifestMessageHeader); i { + switch v := v.(*TruncateCollectionMessageBody); i { case 0: return &v.state case 1: @@ -9168,7 +9244,7 @@ func file_messages_proto_init() { } } file_messages_proto_msgTypes[116].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*BatchUpdateManifestMessageBody); i { + switch v := v.(*BatchUpdateManifestMessageHeader); i { case 0: return &v.state case 1: @@ -9180,7 +9256,7 @@ func file_messages_proto_init() { } } file_messages_proto_msgTypes[117].Exporter = func(v interface{}, i int) interface{} { - switch v := v.(*BatchUpdateManifestItem); i { + switch v := v.(*BatchUpdateManifestMessageBody); i { case 0: return &v.state case 1: @@ -9192,6 +9268,18 @@ func file_messages_proto_init() { } } file_messages_proto_msgTypes[118].Exporter = func(v interface{}, i int) interface{} { + switch v := v.(*BatchUpdateManifestItem); i { + case 0: + return &v.state + case 1: + return &v.sizeCache + case 2: + return &v.unknownFields + default: + return nil + } + } + file_messages_proto_msgTypes[119].Exporter = func(v interface{}, i int) interface{} { switch v := v.(*BatchUpdateManifestV2ColumnGroups); i { case 0: return &v.state @@ -9214,7 +9302,7 @@ func file_messages_proto_init() { GoPackagePath: reflect.TypeOf(x{}).PkgPath(), RawDescriptor: file_messages_proto_rawDesc, NumEnums: 3, - NumMessages: 124, + NumMessages: 125, NumExtensions: 0, NumServices: 0, }, diff --git a/pkg/streaming/util/message/adaptor/message.go b/pkg/streaming/util/message/adaptor/message.go index 40d55936d8..80e1831f76 100644 --- a/pkg/streaming/util/message/adaptor/message.go +++ b/pkg/streaming/util/message/adaptor/message.go @@ -1,6 +1,8 @@ package adaptor import ( + "context" + "github.com/cockroachdb/errors" "github.com/milvus-io/milvus-proto/go-api/v3/commonpb" @@ -78,13 +80,15 @@ func parseTxnMsg(msg message.ImmutableMessage) ([]msgstream.TsMsg, error) { panic("unreachable code, message must be a txn message") } + txnTraceCtx := message.ExtractTraceContext(context.Background(), msg) tsMsgs := make([]msgstream.TsMsg, 0, txnMsg.Size()) err := txnMsg.RangeOver(func(im message.ImmutableMessage) error { var tsMsg msgstream.TsMsg - tsMsg, err := parseSingleMsg(im) + tsMsg, err := parseSingleMsgPayload(im) if err != nil { return err } + tsMsg.SetTraceCtx(txnTraceCtx) tsMsgs = append(tsMsgs, tsMsg) return nil }) @@ -96,14 +100,29 @@ func parseTxnMsg(msg message.ImmutableMessage) ([]msgstream.TsMsg, error) { // parseSingleMsg converts message to ts message. func parseSingleMsg(msg message.ImmutableMessage) (msgstream.TsMsg, error) { + tsMsg, err := parseSingleMsgPayload(msg) + if err != nil { + return nil, err + } + tsMsg.SetTraceCtx(message.ExtractTraceContext(context.Background(), msg)) + return tsMsg, nil +} + +func parseSingleMsgPayload(msg message.ImmutableMessage) (msgstream.TsMsg, error) { + var tsMsg msgstream.TsMsg + var err error switch msg.Version() { case message.VersionV1, message.VersionOld: - return fromMessageToTsMsgV1(msg) + tsMsg, err = fromMessageToTsMsgV1(msg) case message.VersionV2: - return fromMessageToTsMsgV2(msg) + tsMsg, err = fromMessageToTsMsgV2(msg) default: panic("unsupported message version") } + if err != nil { + return nil, err + } + return tsMsg, nil } // fromMessageToTsMsgV1 converts message to ts message. diff --git a/pkg/streaming/util/message/adaptor/message_test.go b/pkg/streaming/util/message/adaptor/message_test.go index f35071ffbf..7ca6f94652 100644 --- a/pkg/streaming/util/message/adaptor/message_test.go +++ b/pkg/streaming/util/message/adaptor/message_test.go @@ -1,10 +1,13 @@ package adaptor import ( + "context" "testing" "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/trace" "github.com/milvus-io/milvus/pkg/v3/mq/msgstream" "github.com/milvus-io/milvus/pkg/v3/streaming/util/message" @@ -57,6 +60,84 @@ func TestNewMsgPackFromInsertMessage(t *testing.T) { } } +func TestNewMsgPackFromMessageRestoresTraceContext(t *testing.T) { + traceID, err := trace.TraceIDFromHex("0102030405060708090a0b0c0d0e0f10") + require.NoError(t, err) + spanID, err := trace.SpanIDFromHex("0102030405060708") + require.NoError(t, err) + ctx := trace.ContextWithSpanContext(context.Background(), trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: traceID, + SpanID: spanID, + })) + + id := rmq.NewRmqID(1) + mutableMsg := message.CreateTestInsertMessage(t, 3, 1, uint64(time.Now().UnixNano()), id) + message.InjectTraceContext(ctx, mutableMsg) + immutableMsg := mutableMsg.WithOldVersion().IntoImmutableMessage(id) + + pack, err := NewMsgPackFromMessage(immutableMsg) + require.NoError(t, err) + require.NotNil(t, pack) + require.Len(t, pack.Msgs, 1) + assert.Equal(t, traceID, trace.SpanContextFromContext(pack.Msgs[0].TraceCtx()).TraceID()) +} + +func TestNewMsgPackFromTxnMessageUsesTxnTraceContext(t *testing.T) { + txnTraceID, err := trace.TraceIDFromHex("0102030405060708090a0b0c0d0e0f10") + require.NoError(t, err) + txnSpanID, err := trace.SpanIDFromHex("0102030405060708") + require.NoError(t, err) + txnCtx := trace.ContextWithSpanContext(context.Background(), trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: txnTraceID, + SpanID: txnSpanID, + })) + + bodyTraceID, err := trace.TraceIDFromHex("1112131415161718191a1b1c1d1e1f20") + require.NoError(t, err) + bodySpanID, err := trace.SpanIDFromHex("1112131415161718") + require.NoError(t, err) + bodyCtx := trace.ContextWithSpanContext(context.Background(), trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: bodyTraceID, + SpanID: bodySpanID, + })) + + id := rmq.NewRmqID(1) + txnContext := message.TxnContext{TxnID: 1} + begin := message.NewBeginTxnMessageBuilderV2(). + WithVChannel("v1"). + WithHeader(&message.BeginTxnMessageHeader{}). + WithBody(&message.BeginTxnMessageBody{}). + MustBuildMutable(). + WithTxnContext(txnContext). + WithTimeTick(1). + WithLastConfirmed(id). + IntoImmutableMessage(id) + builder := message.NewImmutableTxnMessageBuilder(message.MustAsImmutableBeginTxnMessageV2(begin)) + + body := message.CreateTestInsertMessage(t, 3, 1, 2, id).WithTxnContext(txnContext) + message.InjectTraceContext(bodyCtx, body) + builder.Add(body.IntoImmutableMessage(id)) + + commit := message.NewCommitTxnMessageBuilderV2(). + WithVChannel("v1"). + WithHeader(&message.CommitTxnMessageHeader{}). + WithBody(&message.CommitTxnMessageBody{}). + MustBuildMutable(). + WithTxnContext(txnContext). + WithTimeTick(3). + WithLastConfirmed(id) + message.InjectTraceContext(txnCtx, commit) + txnMsg, err := builder.Build(message.MustAsImmutableCommitTxnMessageV2(commit.IntoImmutableMessage(id))) + require.NoError(t, err) + + pack, err := NewMsgPackFromMessage(txnMsg) + require.NoError(t, err) + require.NotNil(t, pack) + require.Len(t, pack.Msgs, 1) + assert.Equal(t, txnTraceID, trace.SpanContextFromContext(pack.Msgs[0].TraceCtx()).TraceID()) + assert.Equal(t, txnSpanID, trace.SpanContextFromContext(pack.Msgs[0].TraceCtx()).SpanID()) +} + func TestNewMsgPackFromCreateCollectionMessage(t *testing.T) { id := rmq.NewRmqID(1) diff --git a/pkg/streaming/util/message/builder.go b/pkg/streaming/util/message/builder.go index c101237cdc..8aaf50ed6c 100644 --- a/pkg/streaming/util/message/builder.go +++ b/pkg/streaming/util/message/builder.go @@ -464,6 +464,10 @@ func newImmutableTxnMesasgeFromWAL( WithTxnContext(*commit.TxnContext()). WithReplicateHeader(commit.ReplicateHeader()). IntoImmutableMessage(commit.MessageID()) + // The assembled txn message uses CommitTxn's trace as the txn-level trace. + if traceContext, ok := commit.Properties().Get(messageTraceContext); ok { + immutableMessage.(*immutableMessageImpl).properties.Set(messageTraceContext, traceContext) + } return &immutableTxnMessageImpl{ immutableMessageImpl: *immutableMessage.(*immutableMessageImpl), begin: MustAsImmutableBeginTxnMessageV2(beginImmutable), diff --git a/pkg/streaming/util/message/message_impl.go b/pkg/streaming/util/message/message_impl.go index b7e1c4b0dd..f0c2bb9b5e 100644 --- a/pkg/streaming/util/message/message_impl.go +++ b/pkg/streaming/util/message/message_impl.go @@ -1,6 +1,7 @@ package message import ( + "context" "fmt" "github.com/milvus-io/milvus-proto/go-api/v3/commonpb" @@ -106,6 +107,18 @@ func (m *messageImpl) WithWALTerm(term int64) MutableMessage { return m } +func (m *messageImpl) injectTraceContext(ctx context.Context) { + injectTraceContext(ctx, m.properties) +} + +func (m *messageImpl) overwriteTraceContext(ctx context.Context) { + overwriteTraceContext(ctx, m.properties) +} + +func (m *immutableMessageImpl) overwriteTraceContext(ctx context.Context) { + overwriteTraceContext(ctx, m.properties) +} + // WithReplicateHeader sets the replicate header of current message. func (m *messageImpl) WithReplicateHeader(rh *ReplicateHeader) MutableMessage { if rh == nil { diff --git a/pkg/streaming/util/message/properties.go b/pkg/streaming/util/message/properties.go index 5e3cc84cd1..c4d458c305 100644 --- a/pkg/streaming/util/message/properties.go +++ b/pkg/streaming/util/message/properties.go @@ -17,6 +17,7 @@ const ( messageNotPersisteted = "_np" // check if the message is unpersisted. messagePChannelLevel = "_pcl" // mark the message as pchannel level message. messageReplicateMesssageHeader = "_rh" // replicate message header. + messageTraceContext = "_tc" // Trace context subset header. ) var ( diff --git a/pkg/streaming/util/message/trace.go b/pkg/streaming/util/message/trace.go new file mode 100644 index 0000000000..ab1e5b27f0 --- /dev/null +++ b/pkg/streaming/util/message/trace.go @@ -0,0 +1,201 @@ +package message + +import ( + "context" + + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/trace" + "go.opentelemetry.io/otel/trace/noop" + + "github.com/milvus-io/milvus/pkg/v3/proto/messagespb" +) + +var noopSpan trace.Span = noop.Span{} + +const ( + tracerName = "milvus.streaming.wal" + + spanAttrMessageType = "message.type" + spanAttrVChannel = "message.vchannel" + spanAttrTimeTick = "message.timetick" + spanAttrReplicate = "message.replicate" + spanAttrTxnID = "txn.id" + + spanAttrBroadcastID = "broadcast.id" + spanAttrBroadcastVChannels = "broadcast.vchannels" + + SpanNameWALAutocommit = "wal.autocommit" + SpanNameWALTxn = "wal.txn" + SpanNameWALBroadcast = "wal.broadcast" + SpanNameWALAppend = "wal.append" + SpanNameWALAppendImpl = "wal.appendimpl" + SpanNameWALDistAppend = "wal.dist_append" + SpanNameWALCatchupConsume = "wal.catchup_consume" + SpanNameWALDistConsume = "wal.dist_consume" + SpanNameReplicateSecondary = "replicate.secondary" + SpanNameWALBCCallback = "wal.bc_callback" +) + +func StartSpan(ctx context.Context, spanName string) (context.Context, trace.Span) { + return otel.Tracer(tracerName).Start(ctx, spanName) +} + +func StartSpanForMessage(ctx context.Context, msg BasicMessage, spanName string) (context.Context, trace.Span) { + if !shouldTraceMessage(msg) { + return ctx, noopSpan + } + ctx, span := otel.Tracer(tracerName).Start(ctx, spanName) + if !span.IsRecording() { + return ctx, span + } + span.SetAttributes(buildMessageSpanAttributes(msg)...) + return ctx, span +} + +func buildMessageSpanAttributes(msg BasicMessage) []attribute.KeyValue { + attrs := []attribute.KeyValue{ + attribute.String(spanAttrMessageType, msg.MessageType().String()), + attribute.String(spanAttrVChannel, getVChannel(msg)), + attribute.Bool(spanAttrReplicate, isReplicateMessage(msg)), + } + if msg.Properties().Exist(messageTimeTick) { + attrs = append(attrs, attribute.Int64(spanAttrTimeTick, int64(msg.TimeTick()))) + } + if txnCtx := msg.TxnContext(); txnCtx != nil { + attrs = append(attrs, attribute.Int64(spanAttrTxnID, int64(txnCtx.TxnID))) + } + if broadcastHeader := msg.BroadcastHeader(); broadcastHeader != nil { + attrs = append(attrs, + attribute.Int64(spanAttrBroadcastID, int64(broadcastHeader.BroadcastID)), + attribute.StringSlice(spanAttrBroadcastVChannels, broadcastHeader.VChannels), + ) + } + return attrs +} + +func shouldTraceMessage(msg BasicMessage) bool { + return msg != nil && msg.MessageType() != MessageTypeTimeTick +} + +func getVChannel(msg BasicMessage) string { + if vchannel, ok := msg.Properties().Get(messageVChannel); ok { + return vchannel + } + return "" +} + +func isReplicateMessage(msg BasicMessage) bool { + return msg.Properties().Exist(messageReplicateMesssageHeader) +} + +type traceContextInjector interface { + injectTraceContext(context.Context) +} + +type traceContextOverwriter interface { + overwriteTraceContext(context.Context) +} + +// InjectTraceContext writes the current span context subset into msg under the +// reserved key _tc as a base64-encoded marshaled TraceContextHeader. +// No-op when _tc already exists or no active / valid span is present on ctx. +// The caller must exclusively own msg because injection mutates Properties. +func InjectTraceContext(ctx context.Context, msg BasicMessage) { + if !shouldTraceMessage(msg) { + return + } + if writer, ok := msg.(traceContextInjector); ok { + writer.injectTraceContext(ctx) + } +} + +// OverwriteTraceContext writes the current span context subset into msg under +// the reserved key _tc even when a trace context already exists. +// The caller must exclusively own msg because overwrite mutates Properties. +func OverwriteTraceContext(ctx context.Context, msg BasicMessage) { + if !shouldTraceMessage(msg) { + return + } + if writer, ok := msg.(traceContextOverwriter); ok { + writer.overwriteTraceContext(ctx) + } +} + +func injectTraceContext(ctx context.Context, p Properties) { + if p.Exist(messageTraceContext) { + return + } + overwriteTraceContext(ctx, p) +} + +func overwriteTraceContext(ctx context.Context, p Properties) { + sc := trace.SpanContextFromContext(ctx) + if !sc.IsValid() { + return + } + val, ok := encodeTraceContextHeader(sc) + if !ok { + return + } + p.Set(messageTraceContext, val) +} + +// ExtractTraceContext reads _tc from msg and returns ctx with the extracted +// remote span context attached. Returns ctx unchanged when _tc is absent or +// malformed — trace propagation is never a correctness dependency. +func ExtractTraceContext(ctx context.Context, msg BasicMessage) context.Context { + sc := extractSpanContext(msg) + if !sc.IsValid() { + return ctx + } + return trace.ContextWithRemoteSpanContext(ctx, sc) +} + +func extractSpanContext(msg BasicMessage) trace.SpanContext { + if msg == nil { + return trace.SpanContext{} + } + return extractSpanContextFromProperties(msg.Properties()) +} + +func extractSpanContextFromProperties(p RProperties) trace.SpanContext { + value, ok := p.Get(messageTraceContext) + if !ok { + return trace.SpanContext{} + } + hdr := &messagespb.TraceContextHeader{} + if err := DecodeProto(value, hdr); err != nil { + return trace.SpanContext{} + } + if len(hdr.GetTraceId()) != 16 || len(hdr.GetSpanId()) != 8 { + return trace.SpanContext{} + } + var tid trace.TraceID + var sid trace.SpanID + copy(tid[:], hdr.GetTraceId()) + copy(sid[:], hdr.GetSpanId()) + return trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: tid, + SpanID: sid, + TraceFlags: trace.TraceFlags(hdr.GetFlags()), + Remote: true, + }) +} + +// encodeTraceContextHeader returns the base64-encoded TraceContextHeader for +// the given span context subset. ok=false when the proto marshal fails. +func encodeTraceContextHeader(sc trace.SpanContext) (string, bool) { + tid := sc.TraceID() + sid := sc.SpanID() + hdr := &messagespb.TraceContextHeader{ + TraceId: tid[:], + SpanId: sid[:], + Flags: uint32(sc.TraceFlags()), + } + val, err := EncodeProto(hdr) + if err != nil { + return "", false + } + return val, true +} diff --git a/pkg/streaming/util/message/trace_test.go b/pkg/streaming/util/message/trace_test.go new file mode 100644 index 0000000000..f6fd790474 --- /dev/null +++ b/pkg/streaming/util/message/trace_test.go @@ -0,0 +1,386 @@ +package message + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" + "go.opentelemetry.io/otel/trace" + + "github.com/milvus-io/milvus-proto/go-api/v3/commonpb" + "github.com/milvus-io/milvus-proto/go-api/v3/msgpb" + "github.com/milvus-io/milvus/pkg/v3/proto/messagespb" +) + +func makeSpanContext(t *testing.T, tid, sid string, flags byte) trace.SpanContext { + t.Helper() + traceID, err := trace.TraceIDFromHex(tid) + assert.NoError(t, err) + spanID, err := trace.SpanIDFromHex(sid) + assert.NoError(t, err) + return trace.NewSpanContext(trace.SpanContextConfig{ + TraceID: traceID, + SpanID: spanID, + TraceFlags: trace.TraceFlags(flags), + Remote: true, + }) +} + +func TestInjectExtractTraceContext_RoundTrip(t *testing.T) { + sc := makeSpanContext(t, "0102030405060708090a0b0c0d0e0f10", "1112131415161718", 0x01) + ctxIn := trace.ContextWithRemoteSpanContext(context.Background(), sc) + + msg := CreateTestEmptyInsertMesage(1, nil) + InjectTraceContext(ctxIn, msg) + + _, ok := msg.Properties().Get(messageTraceContext) + assert.True(t, ok, "InjectTraceContext should set _tc") + + ctxOut := ExtractTraceContext(context.Background(), msg) + got := trace.SpanContextFromContext(ctxOut) + assert.True(t, got.IsValid()) + assert.Equal(t, sc.TraceID(), got.TraceID()) + assert.Equal(t, sc.SpanID(), got.SpanID()) + assert.Equal(t, sc.TraceFlags(), got.TraceFlags()) +} + +func TestInjectTraceContext_NoActiveSpan_NoOp(t *testing.T) { + msg := CreateTestEmptyInsertMesage(1, nil) + InjectTraceContext(context.Background(), msg) // no span in ctx + _, ok := msg.Properties().Get(messageTraceContext) + assert.False(t, ok, "InjectTraceContext should be a no-op without an active span") +} + +func TestInjectTraceContext_ExistingTraceContext_NoOp(t *testing.T) { + existingSC := makeSpanContext(t, "0102030405060708090a0b0c0d0e0f10", "1112131415161718", 0x01) + nextSC := makeSpanContext(t, "2122232425262728292a2b2c2d2e2f30", "3132333435363738", 0x01) + + msg := CreateTestEmptyInsertMesage(1, nil) + InjectTraceContext(trace.ContextWithRemoteSpanContext(context.Background(), existingSC), msg) + InjectTraceContext(trace.ContextWithRemoteSpanContext(context.Background(), nextSC), msg) + + got := trace.SpanContextFromContext(ExtractTraceContext(context.Background(), msg)) + assert.True(t, got.IsValid()) + assert.Equal(t, existingSC.TraceID(), got.TraceID()) + assert.Equal(t, existingSC.SpanID(), got.SpanID()) + assert.Equal(t, existingSC.TraceFlags(), got.TraceFlags()) +} + +func TestOverwriteTraceContext_ExistingTraceContext_Replaces(t *testing.T) { + existingSC := makeSpanContext(t, "0102030405060708090a0b0c0d0e0f10", "1112131415161718", 0x01) + nextSC := makeSpanContext(t, "2122232425262728292a2b2c2d2e2f30", "3132333435363738", 0x01) + + msg := CreateTestEmptyInsertMesage(1, nil) + InjectTraceContext(trace.ContextWithRemoteSpanContext(context.Background(), existingSC), msg) + OverwriteTraceContext(trace.ContextWithRemoteSpanContext(context.Background(), nextSC), msg) + + got := trace.SpanContextFromContext(ExtractTraceContext(context.Background(), msg)) + assert.True(t, got.IsValid()) + assert.Equal(t, nextSC.TraceID(), got.TraceID()) + assert.Equal(t, nextSC.SpanID(), got.SpanID()) + assert.Equal(t, nextSC.TraceFlags(), got.TraceFlags()) +} + +func TestTraceContext_TimeTickMessage_NoOp(t *testing.T) { + sc := makeSpanContext(t, "0102030405060708090a0b0c0d0e0f10", "1112131415161718", 0x01) + ctx := trace.ContextWithRemoteSpanContext(context.Background(), sc) + + msg, err := NewTimeTickMessageBuilderV1(). + WithHeader(&TimeTickMessageHeader{}). + WithBody(&msgpb.TimeTickMsg{}). + WithAllVChannel(). + BuildMutable() + assert.NoError(t, err) + msg.WithTimeTick(100).WithLastConfirmedUseMessageID() + spanCtx, span := StartSpanForMessage(ctx, msg, SpanNameWALAppend) + assert.Equal(t, ctx, spanCtx) + assert.NotNil(t, span) + assert.False(t, span.IsRecording()) + span.RecordError(assert.AnError) + span.SetStatus(0, "") + span.End() + + InjectTraceContext(ctx, msg) + OverwriteTraceContext(ctx, msg) + _, ok := msg.Properties().Get(messageTraceContext) + assert.False(t, ok) +} + +func TestStartSpanForMessage_AddsMessageAttributes(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + msg := CreateTestEmptyInsertMesage(1, nil) + msg.WithTimeTick(100).WithTxnContext(TxnContext{TxnID: 42}) + _, span := StartSpanForMessage(context.Background(), msg, SpanNameWALAppend) + span.End() + + spans := exporter.GetSpans() + require.Len(t, spans, 1) + assert.Equal(t, SpanNameWALAppend, spans[0].Name) + assertSpanAttribute(t, spans[0].Attributes, spanAttrMessageType, MessageTypeInsert.String()) + assertSpanAttribute(t, spans[0].Attributes, spanAttrVChannel, "v1") + assertSpanInt64Attribute(t, spans[0].Attributes, spanAttrTimeTick, 100) + assertSpanInt64Attribute(t, spans[0].Attributes, spanAttrTxnID, 42) + assertSpanBoolAttribute(t, spans[0].Attributes, spanAttrReplicate, false) + + msgID := testMessageID("1") + msg.WithReplicateHeader(&ReplicateHeader{ + ClusterID: "cluster", + MessageID: msgID, + LastConfirmedMessageID: msgID, + TimeTick: 100, + VChannel: "v1", + }) + _, span = StartSpanForMessage(context.Background(), msg, SpanNameWALAppend) + span.End() + + spans = exporter.GetSpans() + require.Len(t, spans, 2) + assertSpanBoolAttribute(t, spans[1].Attributes, spanAttrReplicate, true) +} + +func TestStartSpanForMessage_AddsBroadcastAttributes(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + msg := NewDropCollectionMessageBuilderV1(). + WithHeader(&messagespb.DropCollectionMessageHeader{}). + WithBody(&msgpb.DropCollectionRequest{}). + WithBroadcast([]string{"v1", "v2"}). + MustBuildBroadcast(). + OverwriteBroadcastHeader(11) + _, span := StartSpanForMessage(context.Background(), msg, SpanNameWALBroadcast) + span.End() + + spans := exporter.GetSpans() + require.Len(t, spans, 1) + assert.Equal(t, SpanNameWALBroadcast, spans[0].Name) + assertSpanAttribute(t, spans[0].Attributes, spanAttrMessageType, MessageTypeDropCollection.String()) + assertSpanInt64Attribute(t, spans[0].Attributes, spanAttrBroadcastID, 11) + assertSpanStringSliceAttribute(t, spans[0].Attributes, spanAttrBroadcastVChannels, []string{"v1", "v2"}) +} + +func TestStartSpanForMessage_SkipsAttributesWhenNotRecording(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.NeverSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + msg := CreateTestEmptyInsertMesage(1, nil) + msg.Properties().ToRawMap()[messageTimeTick] = "dirty-timetick" + + assert.NotPanics(t, func() { + _, span := StartSpanForMessage(context.Background(), msg, SpanNameWALAppend) + assert.False(t, span.IsRecording()) + span.End() + }) + assert.Empty(t, exporter.GetSpans()) +} + +func TestImmutableTxnMessageBuildCopiesCommitTraceContext(t *testing.T) { + exporter := tracetest.NewInMemoryExporter() + tp := sdktrace.NewTracerProvider( + sdktrace.WithSyncer(exporter), + sdktrace.WithSampler(sdktrace.AlwaysSample()), + ) + prev := otel.GetTracerProvider() + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(prev) + + sourceCtx, sourceSpan := otel.Tracer("test").Start(context.Background(), SpanNameWALCatchupConsume) + sourceSC := trace.SpanContextFromContext(sourceCtx) + sourceSpan.End() + + txnCtx := TxnContext{TxnID: 1} + lastConfirmed := testMessageID("1") + beginID := testMessageID("2") + begin := NewBeginTxnMessageBuilderV2(). + WithVChannel("v1"). + WithHeader(&BeginTxnMessageHeader{}). + WithBody(&BeginTxnMessageBody{}). + MustBuildMutable(). + WithTxnContext(txnCtx). + WithTimeTick(100). + WithLastConfirmed(lastConfirmed). + IntoImmutableMessage(beginID) + builder := NewImmutableTxnMessageBuilder(MustAsImmutableBeginTxnMessageV2(begin)) + + bodyID := testMessageID("3") + body := CreateTestEmptyInsertMesage(1, nil). + WithTxnContext(txnCtx). + WithTimeTick(101). + WithLastConfirmed(lastConfirmed). + IntoImmutableMessage(bodyID) + builder.Add(body) + + commitID := testMessageID("4") + commit := NewCommitTxnMessageBuilderV2(). + WithVChannel("v1"). + WithHeader(&CommitTxnMessageHeader{}). + WithBody(&CommitTxnMessageBody{}). + MustBuildMutable(). + WithTxnContext(txnCtx). + WithTimeTick(102). + WithLastConfirmed(lastConfirmed) + InjectTraceContext(sourceCtx, commit) + + txn, err := builder.Build(MustAsImmutableCommitTxnMessageV2(commit.IntoImmutableMessage(commitID))) + require.NoError(t, err) + + txnSC := trace.SpanContextFromContext(ExtractTraceContext(context.Background(), txn)) + assert.Equal(t, sourceSC.TraceID(), txnSC.TraceID()) + assert.Equal(t, sourceSC.SpanID(), txnSC.SpanID()) +} + +func TestExtractTraceContext_MissingKey_ReturnsOriginalCtx(t *testing.T) { + msg := CreateTestEmptyInsertMesage(1, nil) + baseCtx := context.Background() + ctx := ExtractTraceContext(baseCtx, msg) + assert.Equal(t, baseCtx, ctx) +} + +func assertSpanAttribute(t *testing.T, attrs []attribute.KeyValue, key string, value string) { + t.Helper() + for _, attr := range attrs { + if string(attr.Key) == key { + assert.Equal(t, value, attr.Value.AsString()) + return + } + } + t.Fatalf("missing span attribute %q", key) +} + +func assertSpanBoolAttribute(t *testing.T, attrs []attribute.KeyValue, key string, value bool) { + t.Helper() + for _, attr := range attrs { + if string(attr.Key) == key { + assert.Equal(t, value, attr.Value.AsBool()) + return + } + } + t.Fatalf("missing span attribute %q", key) +} + +func assertSpanInt64Attribute(t *testing.T, attrs []attribute.KeyValue, key string, value int64) { + t.Helper() + for _, attr := range attrs { + if string(attr.Key) == key { + assert.Equal(t, value, attr.Value.AsInt64()) + return + } + } + t.Fatalf("missing span attribute %q", key) +} + +func assertSpanStringSliceAttribute(t *testing.T, attrs []attribute.KeyValue, key string, value []string) { + t.Helper() + for _, attr := range attrs { + if string(attr.Key) == key { + assert.ElementsMatch(t, value, attr.Value.AsStringSlice()) + return + } + } + t.Fatalf("missing span attribute %q", key) +} + +type testMessageID string + +func (id testMessageID) WALName() WALName { + return WALNameTest +} + +func (id testMessageID) LT(MessageID) bool { + return false +} + +func (id testMessageID) LTE(MessageID) bool { + return true +} + +func (id testMessageID) EQ(other MessageID) bool { + return id.String() == other.String() +} + +func (id testMessageID) Marshal() string { + return string(id) +} + +func (id testMessageID) IntoProto() *commonpb.MessageID { + return &commonpb.MessageID{ + WALName: commonpb.WALName(id.WALName()), + Id: id.Marshal(), + } +} + +func (id testMessageID) String() string { + return string(id) +} + +func TestExtractTraceContext_MalformedValue_ReturnsOriginalCtx(t *testing.T) { + msg := CreateTestEmptyInsertMesage(1, nil) + msg.Properties().ToRawMap()[messageTraceContext] = "!!!not-base64!!!" + baseCtx := context.Background() + ctx := ExtractTraceContext(baseCtx, msg) + assert.Equal(t, baseCtx, ctx) +} + +func TestExtractSpanContext_EmptyWhenMissing(t *testing.T) { + msg := CreateTestEmptyInsertMesage(1, nil) + sc := extractSpanContext(msg) + assert.False(t, sc.IsValid()) +} + +func TestExtractSpanContext_Valid(t *testing.T) { + orig := makeSpanContext(t, "0102030405060708090a0b0c0d0e0f10", "1112131415161718", 0x01) + msg := CreateTestEmptyInsertMesage(1, nil) + InjectTraceContext(trace.ContextWithRemoteSpanContext(context.Background(), orig), msg) + sc := extractSpanContext(msg) + assert.True(t, sc.IsValid()) + assert.Equal(t, orig.TraceID(), sc.TraceID()) + assert.Equal(t, orig.SpanID(), sc.SpanID()) + assert.Equal(t, orig.TraceFlags(), sc.TraceFlags()) +} + +func TestInjectTraceContext_BroadcastMutableMessage(t *testing.T) { + sc := makeSpanContext(t, "0102030405060708090a0b0c0d0e0f10", "1112131415161718", 0x01) + ctx := trace.ContextWithRemoteSpanContext(context.Background(), sc) + + msg := NewDropCollectionMessageBuilderV1(). + WithHeader(&messagespb.DropCollectionMessageHeader{}). + WithBody(&msgpb.DropCollectionRequest{}). + WithBroadcast([]string{"v1", "v2"}). + MustBuildBroadcast() + InjectTraceContext(ctx, msg) + + _, ok := msg.Properties().Get(messageTraceContext) + assert.True(t, ok) + + ctxOut := ExtractTraceContext(context.Background(), msg) + got := trace.SpanContextFromContext(ctxOut) + assert.True(t, got.IsValid()) + assert.Equal(t, sc.TraceID(), got.TraceID()) + assert.Equal(t, sc.SpanID(), got.SpanID()) + assert.Equal(t, sc.TraceFlags(), got.TraceFlags()) +}