Skip to content

Commit 694e2ad

Browse files
committed
fix evaluator
1 parent 5c35c34 commit 694e2ad

3 files changed

Lines changed: 161 additions & 2 deletions

File tree

backend/modules/evaluation/application/experiment_app.go

Lines changed: 74 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -565,8 +565,71 @@ func (e *experimentApplication) SubmitExperiment(ctx context.Context, req *expt.
565565
// 2) 从有序 EvaluatorIDVersionList 中批量解析并按输入顺序回填版本ID
566566
// 3) 从 EvaluatorIDVersionList 中提取 runconfig 和权重配置,构建 evaluator_version_id 到 runconfig/权重的映射
567567
// 注意:runconfig 用于评估器运行时配置,score weight 用于加权分数计算
568+
// validateEvaluatorVersionsBelongToWorkspace 校验直接传入的 evaluator_version_id 是否属于当前工作空间
569+
// 预置评估器(Builtin=true)允许跨空间复用,不做 SpaceID 校验
570+
func (e *experimentApplication) validateEvaluatorVersionsBelongToWorkspace(ctx context.Context, evaluatorVersionIDs []int64, workspaceID int64) error {
571+
if len(evaluatorVersionIDs) == 0 || workspaceID <= 0 {
572+
return nil
573+
}
574+
// 去重,避免重复查询
575+
seen := make(map[int64]struct{}, len(evaluatorVersionIDs))
576+
uniq := make([]int64, 0, len(evaluatorVersionIDs))
577+
for _, id := range evaluatorVersionIDs {
578+
if id <= 0 {
579+
continue
580+
}
581+
if _, ok := seen[id]; ok {
582+
continue
583+
}
584+
seen[id] = struct{}{}
585+
uniq = append(uniq, id)
586+
}
587+
if len(uniq) == 0 {
588+
return nil
589+
}
590+
evs, err := e.evaluatorService.BatchGetEvaluatorVersion(ctx, nil, uniq, false)
591+
if err != nil {
592+
return err
593+
}
594+
found := make(map[int64]*entity.Evaluator, len(evs))
595+
for _, ev := range evs {
596+
if ev == nil {
597+
continue
598+
}
599+
found[ev.GetEvaluatorVersionID()] = ev
600+
}
601+
for _, id := range uniq {
602+
ev, ok := found[id]
603+
if !ok || ev == nil {
604+
return errorx.NewByCode(
605+
errno.EvaluatorVersionNotFoundCode,
606+
errorx.WithExtraMsg(fmt.Sprintf("evaluator version %d not found", id)),
607+
)
608+
}
609+
// 预置评估器允许跨空间复用
610+
if ev.Builtin {
611+
continue
612+
}
613+
if ev.GetSpaceID() != workspaceID {
614+
return errorx.NewByCode(
615+
errno.EvaluatorVersionNotFoundCode,
616+
errorx.WithExtraMsg(fmt.Sprintf("evaluator %d version %s does not belong to workspace %d", ev.ID, ev.GetVersion(), workspaceID)),
617+
)
618+
}
619+
}
620+
return nil
621+
}
622+
568623
func (e *experimentApplication) resolveEvaluatorVersionIDsFromCreateReq(ctx context.Context, req *expt.CreateExperimentRequest) ([]int64, map[int64]*evaluatordto.EvaluatorRunConfig, map[int64]float64, error) {
624+
workspaceID := req.GetWorkspaceID()
625+
569626
evalVersionIDs := make([]int64, 0, len(req.EvaluatorVersionIds))
627+
// 对于直接传入的 evaluator_version_id,需要校验是否属于当前空间(预置评估器除外)
628+
if len(req.EvaluatorVersionIds) > 0 && workspaceID > 0 {
629+
if err := e.validateEvaluatorVersionsBelongToWorkspace(ctx, req.EvaluatorVersionIds, workspaceID); err != nil {
630+
return nil, nil, nil, err
631+
}
632+
}
570633
evalVersionIDs = append(evalVersionIDs, req.EvaluatorVersionIds...)
571634

572635
// 权重映射:key 为 evaluator_version_id,value 为权重(用于加权分数计算)
@@ -609,6 +672,7 @@ func (e *experimentApplication) resolveEvaluatorVersionIDsFromCreateReq(ctx cont
609672
}
610673
for _, ev := range evs {
611674
if ev != nil {
675+
// 预置评估器允许跨空间复用,这里不做 SpaceID 校验
612676
id2Builtin[ev.ID] = ev
613677
}
614678
}
@@ -624,6 +688,16 @@ func (e *experimentApplication) resolveEvaluatorVersionIDsFromCreateReq(ctx cont
624688
if ev == nil {
625689
continue
626690
}
691+
// 非预置评估器必须与实验 WorkspaceID 一致,防止绑定其他空间的评估器
692+
// 同时校验根字段 SpaceID(来自 evaluator 元信息)和内层版本 SpaceID(来自 evaluator_version)
693+
if workspaceID > 0 && !ev.Builtin {
694+
if ev.GetSpaceID() != workspaceID {
695+
return nil, nil, nil, errorx.NewByCode(
696+
errno.EvaluatorVersionNotFoundCode,
697+
errorx.WithExtraMsg(fmt.Sprintf("evaluator %d version %s does not belong to workspace %d", ev.ID, ev.GetVersion(), workspaceID)),
698+
)
699+
}
700+
}
627701
key := fmt.Sprintf("%d#%s", ev.ID, ev.GetVersion())
628702
pair2Eval[key] = ev
629703
}

backend/modules/evaluation/application/experiment_app_test.go

Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -163,6 +163,19 @@ func TestExperimentApplication_CreateExperiment(t *testing.T) {
163163
},
164164
},
165165
mockSetup: func() {
166+
mockEvaluatorService.EXPECT().BatchGetEvaluatorVersion(gomock.Any(), gomock.Any(), []int64{10001}, false).Return([]*entity.Evaluator{
167+
{
168+
ID: 3,
169+
SpaceID: validWorkspaceID,
170+
EvaluatorType: entity.EvaluatorTypePrompt,
171+
PromptEvaluatorVersion: &entity.PromptEvaluatorVersion{
172+
SpaceID: validWorkspaceID,
173+
ID: 10001,
174+
EvaluatorID: 3,
175+
Version: "v1",
176+
},
177+
},
178+
}, nil)
166179
mockEvaluatorService.EXPECT().BatchGetBuiltinEvaluator(gomock.Any(), []int64{1, 1}).Return([]*entity.Evaluator{
167180
{
168181
ID: 1,
@@ -178,8 +191,10 @@ func TestExperimentApplication_CreateExperiment(t *testing.T) {
178191
mockEvaluatorService.EXPECT().BatchGetEvaluatorByIDAndVersion(gomock.Any(), gomock.Any()).Return([]*entity.Evaluator{
179192
{
180193
ID: 2,
194+
SpaceID: validWorkspaceID,
181195
EvaluatorType: entity.EvaluatorTypePrompt,
182196
PromptEvaluatorVersion: &entity.PromptEvaluatorVersion{
197+
SpaceID: validWorkspaceID,
183198
ID: 20200,
184199
EvaluatorID: 2,
185200
Version: "1.0.0",
@@ -234,6 +249,63 @@ func TestExperimentApplication_CreateExperiment(t *testing.T) {
234249
},
235250
wantErr: true,
236251
},
252+
{
253+
name: "cross_workspace_evaluator_id_version_list_rejected",
254+
req: &exptpb.CreateExperimentRequest{
255+
WorkspaceID: validWorkspaceID,
256+
EvaluatorIDVersionList: []*evaluator.EvaluatorIDVersionItem{
257+
{EvaluatorID: gptr.Of(int64(2)), Version: gptr.Of("1.0.0")},
258+
},
259+
},
260+
mockSetup: func() {
261+
mockEvaluatorService.EXPECT().BatchGetEvaluatorByIDAndVersion(gomock.Any(), gomock.Any()).Return([]*entity.Evaluator{
262+
{
263+
ID: 2,
264+
SpaceID: validWorkspaceID + 1, // 其他空间
265+
EvaluatorType: entity.EvaluatorTypePrompt,
266+
Builtin: false,
267+
PromptEvaluatorVersion: &entity.PromptEvaluatorVersion{
268+
SpaceID: validWorkspaceID + 1,
269+
ID: 20200,
270+
EvaluatorID: 2,
271+
Version: "1.0.0",
272+
},
273+
},
274+
}, nil)
275+
},
276+
wantErr: true,
277+
wantCode: errno.EvaluatorVersionNotFoundCode,
278+
},
279+
{
280+
name: "cross_workspace_builtin_evaluator_id_version_list_allowed",
281+
req: &exptpb.CreateExperimentRequest{
282+
WorkspaceID: validWorkspaceID,
283+
EvaluatorIDVersionList: []*evaluator.EvaluatorIDVersionItem{
284+
{EvaluatorID: gptr.Of(int64(9)), Version: gptr.Of("1.0.0")},
285+
},
286+
CreateEvalTargetParam: &eval_target.CreateEvalTargetParam{
287+
EvalTargetType: gptr.Of(domain_eval_target.EvalTargetType_CozeBot),
288+
},
289+
},
290+
mockSetup: func() {
291+
mockEvaluatorService.EXPECT().BatchGetEvaluatorByIDAndVersion(gomock.Any(), gomock.Any()).Return([]*entity.Evaluator{
292+
{
293+
ID: 9,
294+
SpaceID: validWorkspaceID + 2, // 预置评估器通常属于平台空间
295+
EvaluatorType: entity.EvaluatorTypePrompt,
296+
Builtin: true, // 预置评估器允许跨空间复用
297+
PromptEvaluatorVersion: &entity.PromptEvaluatorVersion{
298+
SpaceID: validWorkspaceID + 2,
299+
ID: 90900,
300+
EvaluatorID: 9,
301+
Version: "1.0.0",
302+
},
303+
},
304+
}, nil)
305+
mockManager.EXPECT().CreateExpt(gomock.Any(), gomock.Any(), gomock.Any()).Return(validExpt, nil)
306+
},
307+
wantErr: false,
308+
},
237309
{
238310
name: "skip_missing_evaluators",
239311
req: &exptpb.CreateExperimentRequest{
@@ -248,6 +320,19 @@ func TestExperimentApplication_CreateExperiment(t *testing.T) {
248320
},
249321
},
250322
mockSetup: func() {
323+
mockEvaluatorService.EXPECT().BatchGetEvaluatorVersion(gomock.Any(), gomock.Any(), []int64{10001}, false).Return([]*entity.Evaluator{
324+
{
325+
ID: 3,
326+
SpaceID: validWorkspaceID,
327+
EvaluatorType: entity.EvaluatorTypePrompt,
328+
PromptEvaluatorVersion: &entity.PromptEvaluatorVersion{
329+
SpaceID: validWorkspaceID,
330+
ID: 10001,
331+
EvaluatorID: 3,
332+
Version: "v1",
333+
},
334+
},
335+
}, nil)
251336
mockEvaluatorService.EXPECT().BatchGetBuiltinEvaluator(gomock.Any(), []int64{1}).Return([]*entity.Evaluator{
252337
{
253338
ID: 1,

backend/modules/evaluation/infra/repo/evaluator/mysql/evaluator_tag.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -363,7 +363,7 @@ func (dao *EvaluatorTagDAOImpl) querySourceIDsForCondition(ctx context.Context,
363363
return nil, nil
364364
}
365365
}
366-
if len(restrictTo) <= sourceIDInChunkSize {
366+
if restrictTo == nil || len(restrictTo) <= sourceIDInChunkSize {
367367
return dao.querySourceIDsForConditionOnce(ctx, tagType, langType, condition, restrictTo, opts...)
368368
}
369369
set := make(map[int64]struct{})
@@ -412,7 +412,7 @@ func (dao *EvaluatorTagDAOImpl) sourceIDsForNameLike(ctx context.Context, tagTyp
412412
return nil, nil
413413
}
414414
}
415-
if len(restrictTo) <= sourceIDInChunkSize {
415+
if restrictTo == nil || len(restrictTo) <= sourceIDInChunkSize {
416416
return dao.sourceIDsForNameLikeOnce(ctx, tagType, langType, kw, restrictTo, opts...)
417417
}
418418
set := make(map[int64]struct{})

0 commit comments

Comments
 (0)