diff --git a/backend/modules/observability/domain/component/config/config.go b/backend/modules/observability/domain/component/config/config.go index 82e5c6a084..c85000e4d3 100644 --- a/backend/modules/observability/domain/component/config/config.go +++ b/backend/modules/observability/domain/component/config/config.go @@ -247,7 +247,15 @@ func (c *ReflowInsertConfig) GetDatasetInvokeBatchSize(workspaceID int64) int { // TrajectoryMetadataConfig 轨迹 metadata 写入配置 type TrajectoryMetadataConfig struct { // Spaces 按 workspace_id 配置允许写入的 metadata key 规则列表 - Spaces map[int64][]loop_span.MetaKeyRule `mapstructure:"spaces" json:"spaces"` + Spaces map[int64][]loop_span.MetaKeyRule `mapstructure:"spaces" json:"spaces"` + EnableSingleQuery bool `mapstructure:"enable_single_query" json:"enable_single_query"` +} + +func (c *TrajectoryMetadataConfig) IsSingleQueryEnabled() bool { + if c == nil { + return false + } + return c.EnableSingleQuery } //go:generate mockgen -destination=mocks/config.go -package=mocks . ITraceConfig diff --git a/backend/modules/observability/domain/trace/service/trace_service.go b/backend/modules/observability/domain/trace/service/trace_service.go index 1900274022..b9e46910c4 100644 --- a/backend/modules/observability/domain/trace/service/trace_service.go +++ b/backend/modules/observability/domain/trace/service/trace_service.go @@ -2594,6 +2594,66 @@ func (r *TraceServiceImpl) GetTrajectories(ctx context.Context, workspaceID int6 maxBytes = backfillCfg.GetTrajectoryMaxBytes(workspaceID) } + if metaCfg.IsSingleQueryEnabled() { + return r.getTrajectoriesSingleQuery(ctx, workspaceID, tenant, traceIDs, startTime, endTime, platformType, trajectoryConfig, metaRules, maxBytes) + } + + return r.getTrajectoriesDoubleQuery(ctx, workspaceID, tenant, traceIDs, startTime, endTime, platformType, trajectoryConfig, metaRules, maxBytes) +} + +func (r *TraceServiceImpl) getTrajectoriesSingleQuery(ctx context.Context, workspaceID int64, tenant []string, traceIDs []string, startTime, endTime int64, platformType loop_span.PlatformType, trajectoryConfig *GetTrajectoryConfigResponse, metaRules []loop_span.MetaKeyRule, maxBytes int64) (map[string]*loop_span.Trajectory, error) { + allSpans, err := r.traceRepo.ListSpansRepeat(ctx, &repo.ListSpansParam{ + Tenants: tenant, + Filters: &loop_span.FilterFields{ + FilterFields: []*loop_span.FilterField{ + { + FieldName: "trace_id", + FieldType: loop_span.FieldTypeString, + Values: traceIDs, + QueryType: ptr.Of(loop_span.QueryTypeEnumIn), + }, + }, + }, + StartAt: startTime, + EndAt: endTime, + Limit: 1000, + NotQueryAnnotation: true, + MaxBytes: maxBytes, + }) + if err != nil { + if errors.Is(err, repo.ErrMaxBytesExceeded) { + logs.CtxWarn(ctx, "getTrajectoriesSingleQuery skipped: allSpans exceeded max bytes, traceIDs:%v, maxBytes:%d", traceIDs, maxBytes) + return map[string]*loop_span.Trajectory{}, nil + } + logs.CtxError(ctx, "Failed to list all spans, err:%+v", err) + return nil, err + } + + selectFilters := r.getSelectFilters(traceIDs, trajectoryConfig, allSpans) + + processors, err := r.buildHelper.BuildGetTraceProcessors(ctx, span_processor.Settings{ + WorkspaceId: workspaceID, + PlatformType: platformType, + QueryStartTime: startTime, + QueryEndTime: endTime, + SpanDoubleCheck: false, + }) + if err != nil { + return nil, errorx.WrapByCode(err, obErrorx.CommercialCommonInternalErrorCodeCode) + } + allSpans.Spans, err = runSpanProcessors(ctx, "GetTrajectories", processors, allSpans.Spans) + if err != nil { + return nil, err + } + + trajectories, err := r.buildTrajectories(ctx, &allSpans.Spans, ptr.Of(r.convertCustomNode(allSpans.Spans)), selectFilters, metaRules) + if err != nil { + return nil, err + } + return trajectories, nil +} + +func (r *TraceServiceImpl) getTrajectoriesDoubleQuery(ctx context.Context, workspaceID int64, tenant []string, traceIDs []string, startTime, endTime int64, platformType loop_span.PlatformType, trajectoryConfig *GetTrajectoryConfigResponse, metaRules []loop_span.MetaKeyRule, maxBytes int64) (map[string]*loop_span.Trajectory, error) { allSpans, err := r.traceRepo.ListSpansRepeat(ctx, &repo.ListSpansParam{ Tenants: tenant, Filters: &loop_span.FilterFields{ diff --git a/backend/modules/observability/domain/trace/service/trace_trajectory_service_test.go b/backend/modules/observability/domain/trace/service/trace_trajectory_service_test.go index 65673f63e4..e189c1b0cd 100644 --- a/backend/modules/observability/domain/trace/service/trace_trajectory_service_test.go +++ b/backend/modules/observability/domain/trace/service/trace_trajectory_service_test.go @@ -5,9 +5,11 @@ package service import ( "context" + "errors" "testing" "time" + config "github.com/coze-dev/coze-loop/backend/modules/observability/domain/component/config" configmocks "github.com/coze-dev/coze-loop/backend/modules/observability/domain/component/config/mocks" tenantmocks "github.com/coze-dev/coze-loop/backend/modules/observability/domain/component/tenant/mocks" "github.com/coze-dev/coze-loop/backend/modules/observability/domain/trace/entity" @@ -21,6 +23,15 @@ import ( "go.uber.org/mock/gomock" ) +type errorGetTraceProcessorsBuildHelper struct { + TraceFilterProcessorBuilder + err error +} + +func (e *errorGetTraceProcessorsBuildHelper) BuildGetTraceProcessors(_ context.Context, _ span_processor.Settings) ([]span_processor.Processor, error) { + return nil, e.err +} + func TestTraceServiceImpl_GetTrajectoryConfig(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() @@ -151,3 +162,117 @@ func TestTraceServiceImpl_GetTrajectories_and_ListTrajectory(t *testing.T) { assert.NoError(t, err) assert.Equal(t, 1, len(lt.Trajectories)) } + +func TestTraceServiceImpl_GetTrajectories_SingleQuery(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + repoMock := repomocks.NewMockITraceRepo(ctrl) + filterFactoryMock := filtermocks.NewMockPlatformFilterFactory(ctrl) + builder := NewTraceFilterProcessorBuilder(filterFactoryMock, map[entity.ProcessorScene][]span_processor.Factory{ + entity.SceneGetTrace: {span_processor.NewCheckProcessorFactory()}, + }) + tenantProviderMock := tenantmocks.NewMockITenantProvider(ctrl) + tenantProviderMock.EXPECT().GetTenantsByPlatformType(gomock.Any(), gomock.Any()).Return([]string{"tenant"}, nil).AnyTimes() + repoMock.EXPECT().GetTrajectoryConfig(gomock.Any(), repo.GetTrajectoryConfigParam{WorkspaceId: 1}).Return(nil, nil).AnyTimes() + traceConfigMock := configmocks.NewMockITraceConfig(ctrl) + traceConfigMock.EXPECT().GetTrajectoryMetadataConfig(gomock.Any()).Return(&config.TrajectoryMetadataConfig{EnableSingleQuery: true}).AnyTimes() + traceConfigMock.EXPECT().GetBackfillConfig(gomock.Any()).Return(nil).AnyTimes() + + svc := &TraceServiceImpl{traceRepo: repoMock, buildHelper: builder, tenantProvider: tenantProviderMock, traceConfig: traceConfigMock} + traceIDs := []string{"tid"} + + listSpansCallCount := 0 + repoMock.EXPECT().ListSpansRepeat(gomock.Any(), gomock.Any()).DoAndReturn(func(_ context.Context, p *repo.ListSpansParam) (*repo.ListSpansResult, error) { + listSpansCallCount++ + assert.Empty(t, p.SelectColumns, "single query should not set SelectColumns") + return &repo.ListSpansResult{Spans: loop_span.SpanList{ + {TraceID: "tid", SpanID: "root", ParentID: "0", WorkspaceID: "1", SpanName: "root", SpanType: "agent"}, + {TraceID: "tid", SpanID: "m", ParentID: "root", WorkspaceID: "1", SpanName: "model", SpanType: "model"}, + }}, nil + }).AnyTimes() + + res, err := svc.GetTrajectories(context.Background(), 1, traceIDs, time.Now().Add(-time.Minute).UnixMilli(), time.Now().UnixMilli(), loop_span.PlatformCozeLoop) + assert.NoError(t, err) + assert.NotNil(t, res["tid"]) + assert.Equal(t, 1, listSpansCallCount, "single query mode should call ListSpansRepeat exactly once") +} + +func TestTrajectoryMetadataConfig_IsSingleQueryEnabled(t *testing.T) { + var nilCfg *config.TrajectoryMetadataConfig + assert.False(t, nilCfg.IsSingleQueryEnabled()) + + assert.False(t, (&config.TrajectoryMetadataConfig{}).IsSingleQueryEnabled()) + + assert.True(t, (&config.TrajectoryMetadataConfig{EnableSingleQuery: true}).IsSingleQueryEnabled()) +} + +func TestTraceServiceImpl_GetTrajectories_SingleQuery_MaxBytesExceeded(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + repoMock := repomocks.NewMockITraceRepo(ctrl) + filterFactoryMock := filtermocks.NewMockPlatformFilterFactory(ctrl) + builder := NewTraceFilterProcessorBuilder(filterFactoryMock, map[entity.ProcessorScene][]span_processor.Factory{ + entity.SceneGetTrace: {span_processor.NewCheckProcessorFactory()}, + }) + tenantProviderMock := tenantmocks.NewMockITenantProvider(ctrl) + tenantProviderMock.EXPECT().GetTenantsByPlatformType(gomock.Any(), gomock.Any()).Return([]string{"tenant"}, nil) + repoMock.EXPECT().GetTrajectoryConfig(gomock.Any(), gomock.Any()).Return(nil, nil) + traceConfigMock := configmocks.NewMockITraceConfig(ctrl) + traceConfigMock.EXPECT().GetTrajectoryMetadataConfig(gomock.Any()).Return(&config.TrajectoryMetadataConfig{EnableSingleQuery: true}) + traceConfigMock.EXPECT().GetBackfillConfig(gomock.Any()).Return(nil) + + svc := &TraceServiceImpl{traceRepo: repoMock, buildHelper: builder, tenantProvider: tenantProviderMock, traceConfig: traceConfigMock} + + repoMock.EXPECT().ListSpansRepeat(gomock.Any(), gomock.Any()).Return(nil, repo.ErrMaxBytesExceeded) + + res, err := svc.GetTrajectories(context.Background(), 1, []string{"tid"}, time.Now().Add(-time.Minute).UnixMilli(), time.Now().UnixMilli(), loop_span.PlatformCozeLoop) + assert.NoError(t, err) + assert.Empty(t, res) +} + +func TestTraceServiceImpl_GetTrajectories_SingleQuery_ListSpansError(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + repoMock := repomocks.NewMockITraceRepo(ctrl) + filterFactoryMock := filtermocks.NewMockPlatformFilterFactory(ctrl) + builder := NewTraceFilterProcessorBuilder(filterFactoryMock, map[entity.ProcessorScene][]span_processor.Factory{ + entity.SceneGetTrace: {span_processor.NewCheckProcessorFactory()}, + }) + tenantProviderMock := tenantmocks.NewMockITenantProvider(ctrl) + tenantProviderMock.EXPECT().GetTenantsByPlatformType(gomock.Any(), gomock.Any()).Return([]string{"tenant"}, nil) + repoMock.EXPECT().GetTrajectoryConfig(gomock.Any(), gomock.Any()).Return(nil, nil) + traceConfigMock := configmocks.NewMockITraceConfig(ctrl) + traceConfigMock.EXPECT().GetTrajectoryMetadataConfig(gomock.Any()).Return(&config.TrajectoryMetadataConfig{EnableSingleQuery: true}) + traceConfigMock.EXPECT().GetBackfillConfig(gomock.Any()).Return(nil) + + svc := &TraceServiceImpl{traceRepo: repoMock, buildHelper: builder, tenantProvider: tenantProviderMock, traceConfig: traceConfigMock} + + repoMock.EXPECT().ListSpansRepeat(gomock.Any(), gomock.Any()).Return(nil, errors.New("ck connection error")) + + res, err := svc.GetTrajectories(context.Background(), 1, []string{"tid"}, time.Now().Add(-time.Minute).UnixMilli(), time.Now().UnixMilli(), loop_span.PlatformCozeLoop) + assert.Error(t, err) + assert.Nil(t, res) +} + +func TestTraceServiceImpl_GetTrajectories_SingleQuery_BuildProcessorsError(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + repoMock := repomocks.NewMockITraceRepo(ctrl) + buildHelper := &errorGetTraceProcessorsBuildHelper{err: errors.New("build processor error")} + tenantProviderMock := tenantmocks.NewMockITenantProvider(ctrl) + tenantProviderMock.EXPECT().GetTenantsByPlatformType(gomock.Any(), gomock.Any()).Return([]string{"tenant"}, nil) + repoMock.EXPECT().GetTrajectoryConfig(gomock.Any(), gomock.Any()).Return(nil, nil) + traceConfigMock := configmocks.NewMockITraceConfig(ctrl) + traceConfigMock.EXPECT().GetTrajectoryMetadataConfig(gomock.Any()).Return(&config.TrajectoryMetadataConfig{EnableSingleQuery: true}) + traceConfigMock.EXPECT().GetBackfillConfig(gomock.Any()).Return(nil) + + svc := &TraceServiceImpl{traceRepo: repoMock, buildHelper: buildHelper, tenantProvider: tenantProviderMock, traceConfig: traceConfigMock} + + repoMock.EXPECT().ListSpansRepeat(gomock.Any(), gomock.Any()).Return(&repo.ListSpansResult{Spans: loop_span.SpanList{ + {TraceID: "tid", SpanID: "root", ParentID: "0", WorkspaceID: "1", SpanName: "root", SpanType: "agent"}, + }}, nil) + + res, err := svc.GetTrajectories(context.Background(), 1, []string{"tid"}, time.Now().Add(-time.Minute).UnixMilli(), time.Now().UnixMilli(), loop_span.PlatformCozeLoop) + assert.Error(t, err) + assert.Nil(t, res) +}