From cd5254a8fe7974904b622755731bfe8b2cf16e9c Mon Sep 17 00:00:00 2001 From: wlh000 Date: Wed, 5 Aug 2026 10:47:39 +0800 Subject: [PATCH] fix: enforce data session ownership (IK6HAL) --- tms/data-proxy/internal/service/data_svc.go | 1 + .../internal/service/data_svc_test.go | 64 +++++++++++++++++++ tms/data-proxy/internal/session/manager.go | 9 +++ .../internal/session/manager_test.go | 50 +++++++-------- 4 files changed, 99 insertions(+), 25 deletions(-) diff --git a/tms/data-proxy/internal/service/data_svc.go b/tms/data-proxy/internal/service/data_svc.go index 51a8b16..fd1c2be 100644 --- a/tms/data-proxy/internal/service/data_svc.go +++ b/tms/data-proxy/internal/service/data_svc.go @@ -165,6 +165,7 @@ func (s *DataService) UploadFragment(ctx context.Context, req *datapb.UploadFrag sess, status, code, humanMsg, err := s.sm.Upload( ctx, + identity.AgentId, req.GetTaskId(), req.GetOffset(), req.GetSize(), diff --git a/tms/data-proxy/internal/service/data_svc_test.go b/tms/data-proxy/internal/service/data_svc_test.go index 7b52e51..4da528e 100644 --- a/tms/data-proxy/internal/service/data_svc_test.go +++ b/tms/data-proxy/internal/service/data_svc_test.go @@ -381,6 +381,70 @@ func TestUploadFragment_SessionExpired(t *testing.T) { } } +func TestUploadFragment_RejectsDifferentAgent(t *testing.T) { + svc, _, mr := newTestService(t) + ctx := context.Background() + payload := bytes.Repeat([]byte{'x'}, 20) + + sess, err := svc.NewDataSession(ctx, &datapb.NewDataSessionRequest{ + Identity: identity(1001), + DataId: 42, + TotalSize: uint64(len(payload)), + Md5: md5Hex(payload), + }) + if err != nil { + t.Fatal(err) + } + + resp, err := svc.UploadFragment(ctx, &datapb.UploadFragmentRequest{ + Identity: identity(2002), + TaskId: sess.TaskId, + Offset: 0, + Size: 8, + FragmentId: 0, + Data: payload[:8], + }) + if err != nil { + t.Fatal(err) + } + if resp.Status != datapb.SessionStatus_SESSION_STATUS_RESTART { + t.Fatalf("status=%v, want RESTART for different agent", resp.Status) + } + if mr.Exists(session.FragmentKey(sess.TaskId)) { + t.Fatal("different agent wrote fragment data") + } +} + +func TestNewDataSession_RejectsResumeByDifferentAgent(t *testing.T) { + svc, _, _ := newTestService(t) + ctx := context.Background() + payload := bytes.Repeat([]byte{'x'}, 20) + + first, err := svc.NewDataSession(ctx, &datapb.NewDataSessionRequest{ + Identity: identity(1001), + DataId: 42, + TotalSize: uint64(len(payload)), + Md5: md5Hex(payload), + }) + if err != nil { + t.Fatal(err) + } + + resumed, err := svc.NewDataSession(ctx, &datapb.NewDataSessionRequest{ + Identity: identity(2002), + TaskId: first.TaskId, + DataId: 42, + TotalSize: uint64(len(payload)), + Md5: md5Hex(payload), + }) + if err != nil { + t.Fatal(err) + } + if resumed.Status != datapb.SessionStatus_SESSION_STATUS_RESTART { + t.Fatalf("status=%v, want RESTART for different agent", resumed.Status) + } +} + // TestReportError_ReturnsSuccessAlways asserts idempotent fire-and-forget semantics. func TestReportError_ReturnsSuccessAlways(t *testing.T) { svc, _, _ := newTestService(t) diff --git a/tms/data-proxy/internal/session/manager.go b/tms/data-proxy/internal/session/manager.go index f98a770..18b9a9a 100644 --- a/tms/data-proxy/internal/session/manager.go +++ b/tms/data-proxy/internal/session/manager.go @@ -109,6 +109,10 @@ func (m *Manager) CreateOrResume( return nil, datapb.SessionStatus_SESSION_STATUS_RESTART, errcode.CodeInternal, loadErr.Error(), loadErr } + if meta.AgentID != agentID { + return nil, datapb.SessionStatus_SESSION_STATUS_RESTART, + errcode.CodeSessionMismatch, "session owner mismatch", nil + } if meta.DataID != dataID || meta.MD5 != md5Hex { return nil, datapb.SessionStatus_SESSION_STATUS_RESTART, errcode.CodeSessionMismatch, @@ -167,6 +171,7 @@ func (m *Manager) CreateOrResume( // 在末分片正常路径上,则返回会话被最终化时的快照(调用方可能仍需这些 id)。 func (m *Manager) Upload( ctx context.Context, + agentID uint64, taskID uint64, reqOffset uint64, reqSize uint32, @@ -183,6 +188,10 @@ func (m *Manager) Upload( return nil, datapb.SessionStatus_SESSION_STATUS_RESTART, errcode.CodeInternal, loadErr.Error(), loadErr } + if meta.AgentID != agentID { + return nil, datapb.SessionStatus_SESSION_STATUS_RESTART, + errcode.CodeSessionMismatch, "session owner mismatch", nil + } // 60s 窗口内重复的末分片:直接返回 FINISH,不再投递。 if meta.IsFinished() { diff --git a/tms/data-proxy/internal/session/manager_test.go b/tms/data-proxy/internal/session/manager_test.go index cf5da99..4eb7500 100644 --- a/tms/data-proxy/internal/session/manager_test.go +++ b/tms/data-proxy/internal/session/manager_test.go @@ -202,7 +202,7 @@ func TestUpload_NonFinalFragment(t *testing.T) { mgr, _, mr, _, _ := newTestManager(t) sess, payload := createSession(t, mgr, 20) // 3 fragments at fragment_size=8 _, status, code, _, err := mgr.Upload( - context.Background(), sess.TaskID, 0, 8, 0, payload[0:8], + context.Background(), sess.AgentID, sess.TaskID, 0, 8, 0, payload[0:8], ) if err != nil { t.Fatalf("upload err: %v", err) @@ -224,7 +224,7 @@ func TestUpload_DerivesOffsetFromFragmentMetadata(t *testing.T) { sess, payload := createSession(t, mgr, 20) updated, status, code, _, err := mgr.Upload( - context.Background(), sess.TaskID, 1000, 8, 0, payload[0:8], + context.Background(), sess.AgentID, sess.TaskID, 1000, 8, 0, payload[0:8], ) if err != nil { t.Fatalf("upload err: %v", err) @@ -250,9 +250,9 @@ func TestUpload_FinalFragmentTriggersPublish(t *testing.T) { mgr, pub, mr, _, _ := newTestManager(t) sess, payload := createSession(t, mgr, 20) // 8/8/4 ctx := context.Background() - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 0, 8, 0, payload[0:8]) - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 8, 8, 1, payload[8:16]) - _, status, code, _, err := mgr.Upload(ctx, sess.TaskID, 16, 4, 2, payload[16:20]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 0, 8, 0, payload[0:8]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 8, 8, 1, payload[8:16]) + _, status, code, _, err := mgr.Upload(ctx, sess.AgentID, sess.TaskID, 16, 4, 2, payload[16:20]) if err != nil { t.Fatalf("upload final: %v", err) } @@ -278,10 +278,10 @@ func TestUpload_FinalFragmentMD5Mismatch(t *testing.T) { sess, payload := createSession(t, mgr, 20) ctx := context.Background() // Corrupt the last fragment so md5 fails. - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 0, 8, 0, payload[0:8]) - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 8, 8, 1, payload[8:16]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 0, 8, 0, payload[0:8]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 8, 8, 1, payload[8:16]) corrupt := bytes.Repeat([]byte{0xFF}, 4) - _, status, code, _, _ := mgr.Upload(ctx, sess.TaskID, 16, 4, 2, corrupt) + _, status, code, _, _ := mgr.Upload(ctx, sess.AgentID, sess.TaskID, 16, 4, 2, corrupt) if status != datapb.SessionStatus_SESSION_STATUS_RESTART || code != errcode.CodeMD5Mismatch { t.Fatalf("status=%v code=%d", status, code) } @@ -296,9 +296,9 @@ func TestUpload_KafkaPublishFailure(t *testing.T) { sess, payload := createSession(t, mgr, 20) pub.failErr = errors.New("kafka down") ctx := context.Background() - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 0, 8, 0, payload[0:8]) - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 8, 8, 1, payload[8:16]) - _, status, code, _, _ := mgr.Upload(ctx, sess.TaskID, 16, 4, 2, payload[16:20]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 0, 8, 0, payload[0:8]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 8, 8, 1, payload[8:16]) + _, status, code, _, _ := mgr.Upload(ctx, sess.AgentID, sess.TaskID, 16, 4, 2, payload[16:20]) if status != datapb.SessionStatus_SESSION_STATUS_RESTART || code != errcode.CodeInternal { t.Fatalf("status=%v code=%d", status, code) } @@ -308,11 +308,11 @@ func TestUpload_RetryAfterFinishHSetFailureDoesNotRepublish(t *testing.T) { mgr, pub, _, _, rc := newTestManager(t) sess, payload := createSession(t, mgr, 20) ctx := context.Background() - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 0, 8, 0, payload[0:8]) - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 8, 8, 1, payload[8:16]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 0, 8, 0, payload[0:8]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 8, 8, 1, payload[8:16]) mgr.client = &failFinishHSetOnceClient{Cmdable: rc} - _, status, code, _, err := mgr.Upload(ctx, sess.TaskID, 16, 4, 2, payload[16:20]) + _, status, code, _, err := mgr.Upload(ctx, sess.AgentID, sess.TaskID, 16, 4, 2, payload[16:20]) if err == nil || status != datapb.SessionStatus_SESSION_STATUS_RESTART || code != errcode.CodeInternal { t.Fatalf("first final upload: status=%v code=%d err=%v", status, code, err) } @@ -324,7 +324,7 @@ func TestUpload_RetryAfterFinishHSetFailureDoesNotRepublish(t *testing.T) { t.Fatalf("publish receipt = %q, err=%v; want 1", receipt, err) } - _, status, code, _, err = mgr.Upload(ctx, sess.TaskID, 16, 4, 2, payload[16:20]) + _, status, code, _, err = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 16, 4, 2, payload[16:20]) if err != nil || status != datapb.SessionStatus_SESSION_STATUS_FINISH || code != errcode.CodeOK { t.Fatalf("retry final upload: status=%v code=%d err=%v", status, code, err) } @@ -338,7 +338,7 @@ func TestUpload_FragmentOutOfRange(t *testing.T) { mgr, _, _, _, _ := newTestManager(t) sess, _ := createSession(t, mgr, 20) // fragment_count=3 _, status, code, _, _ := mgr.Upload( - context.Background(), sess.TaskID, 0, 8, 5, bytes.Repeat([]byte{'a'}, 8), + context.Background(), sess.AgentID, sess.TaskID, 0, 8, 5, bytes.Repeat([]byte{'a'}, 8), ) if status != datapb.SessionStatus_SESSION_STATUS_RESTART || code != errcode.CodeFragmentOutOfRange { t.Fatalf("status=%v code=%d", status, code) @@ -350,7 +350,7 @@ func TestUpload_FragmentSizeMismatchNonFinal(t *testing.T) { mgr, _, _, _, _ := newTestManager(t) sess, _ := createSession(t, mgr, 20) _, status, code, _, _ := mgr.Upload( - context.Background(), sess.TaskID, 0, 4, 0, bytes.Repeat([]byte{'a'}, 4), + context.Background(), sess.AgentID, sess.TaskID, 0, 4, 0, bytes.Repeat([]byte{'a'}, 4), ) if status != datapb.SessionStatus_SESSION_STATUS_RESTART || code != errcode.CodeFragmentSizeMismatch { t.Fatalf("status=%v code=%d", status, code) @@ -363,14 +363,14 @@ func TestUpload_FinishedReturnsFinishNoRepublish(t *testing.T) { mgr, pub, _, _, _ := newTestManager(t) sess, payload := createSession(t, mgr, 20) ctx := context.Background() - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 0, 8, 0, payload[0:8]) - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 8, 8, 1, payload[8:16]) - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 16, 4, 2, payload[16:20]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 0, 8, 0, payload[0:8]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 8, 8, 1, payload[8:16]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 16, 4, 2, payload[16:20]) if pub.count() != 1 { t.Fatalf("setup: expected 1 publish") } // Duplicate final: must still say FINISH but not republish. - _, status, _, _, _ := mgr.Upload(ctx, sess.TaskID, 16, 4, 2, payload[16:20]) + _, status, _, _, _ := mgr.Upload(ctx, sess.AgentID, sess.TaskID, 16, 4, 2, payload[16:20]) if status != datapb.SessionStatus_SESSION_STATUS_FINISH { t.Errorf("status = %v", status) } @@ -385,11 +385,11 @@ func TestUpload_DuplicateNonFinalIsIdempotent(t *testing.T) { mgr, _, _, _, _ := newTestManager(t) sess, payload := createSession(t, mgr, 20) ctx := context.Background() - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 0, 8, 0, payload[0:8]) - _, _, _, _, _ = mgr.Upload(ctx, sess.TaskID, 8, 8, 1, payload[8:16]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 0, 8, 0, payload[0:8]) + _, _, _, _, _ = mgr.Upload(ctx, sess.AgentID, sess.TaskID, 8, 8, 1, payload[8:16]) // Now re-send fragment 0 (older). Offset/FragmentID MUST NOT regress. - s, status, code, _, _ := mgr.Upload(ctx, sess.TaskID, 0, 8, 0, payload[0:8]) + s, status, code, _, _ := mgr.Upload(ctx, sess.AgentID, sess.TaskID, 0, 8, 0, payload[0:8]) if status != datapb.SessionStatus_SESSION_STATUS_REQUIRE_NEXT || code != errcode.CodeOK { t.Fatalf("status=%v code=%d", status, code) } @@ -402,7 +402,7 @@ func TestUpload_DuplicateNonFinalIsIdempotent(t *testing.T) { func TestUpload_SessionExpired(t *testing.T) { mgr, _, _, _, _ := newTestManager(t) _, status, code, humanMsg, _ := mgr.Upload( - context.Background(), 777, 0, 8, 0, bytes.Repeat([]byte{'a'}, 8), + context.Background(), 10086, 777, 0, 8, 0, bytes.Repeat([]byte{'a'}, 8), ) if status != datapb.SessionStatus_SESSION_STATUS_RESTART || code != errcode.CodeSessionExpired { t.Fatalf("status=%v code=%d", status, code) -- Gitee