diff --git a/arkruntime/lib/environments/heartbeat.go b/arkruntime/lib/environments/heartbeat.go index 88fa267..003a294 100644 --- a/arkruntime/lib/environments/heartbeat.go +++ b/arkruntime/lib/environments/heartbeat.go @@ -7,6 +7,7 @@ import ( "net/http" "time" + "github.com/volcengine/ark-runtime-go/arkruntime/model/environment" selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" ) @@ -25,11 +26,11 @@ const ( heartbeatStopCausePermanentFailure heartbeatStopCause = "heartbeat_permanent_failure" ) -func (w *EnvironmentWorker) heartbeatLoop(ctx context.Context, item selfhosted.WorkItem, api selfhosted.API, cancel context.CancelFunc, markStopped func(heartbeatStopCause)) { +func (w *EnvironmentWorker) heartbeatLoop(ctx context.Context, work claimedWork, api selfhosted.API, cancel context.CancelFunc, markStopped func(heartbeatStopCause)) { interval := clampHeartbeatInterval(heartbeatDefault/2, heartbeatDefault) ttl := heartbeatDefault - logger := w.logger().With("component", "environment-worker", "work_id", item.ID, "session_id", item.SessionIDValue()) - last := item.LatestHeartbeatValue() + logger := w.logger().With("component", "environment-worker", "work_id", work.ID, "session_id", work.SessionID) + last := work.LatestHeartbeatAt if last == "" { last = selfhosted.ExpectedLastHeartbeatNoHeartbeat } @@ -38,10 +39,10 @@ func (w *EnvironmentWorker) heartbeatLoop(ctx context.Context, item selfhosted.W beatCtx, beatCancel := context.WithTimeout(ctx, interval) defer beatCancel() resp, err := api.HeartbeatWork(beatCtx, selfhosted.HeartbeatWorkRequest{ - EnvironmentID: item.EnvironmentID, - WorkID: item.ID, - ExpectedLastHeartbeat: last, - DesiredTTLSeconds: int(ttl / time.Second), + EnvironmentID: work.EnvironmentID, + WorkID: work.ID, + ExpectedLastHeartbeat: environment.NewOptString(last), + DesiredTTLSeconds: environment.NewOptInt64(int64(ttl / time.Second)), }) if err != nil { if selfhosted.IsStatus(err, http.StatusPreconditionFailed) { @@ -95,7 +96,7 @@ func (w *EnvironmentWorker) heartbeatLoop(ctx context.Context, item selfhosted.W cancel() return false } - if resp.LeaseExtended != nil && !*resp.LeaseExtended { + if !resp.LeaseExtended { logger.Warn("heartbeat lease not extended", "state", state) if markStopped != nil { markStopped(heartbeatStopCauseLeaseNotExtended) diff --git a/arkruntime/lib/environments/integration_test.go b/arkruntime/lib/environments/integration_test.go index 44ff59f..cb68359 100644 --- a/arkruntime/lib/environments/integration_test.go +++ b/arkruntime/lib/environments/integration_test.go @@ -50,7 +50,7 @@ func (f *fakeEnvironmentWorkerAPI) HeartbeatWork(ctx context.Context, req selfho } return &selfhosted.HeartbeatResponse{ LastHeartbeat: time.Now().UTC().Format(time.RFC3339Nano), - LeaseExtended: selfhosted.BoolPtr(true), + LeaseExtended: true, State: selfhosted.WorkStateActive, }, nil } @@ -102,12 +102,7 @@ func TestEnvironmentWorkerRunHandlesPolledWorkInProcess(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() api := &fakeEnvironmentWorkerAPI{ - pollItem: &selfhosted.WorkItem{ - ID: "work_local", - EnvironmentID: "env_local", - SessionID: "sess_local", - LeaseID: "lease_local", - }, + pollItem: newTestWorkItem("work_local", "env_local", "sess_local"), events: []selfhosted.Event{ { ID: "toolu_local", @@ -139,10 +134,10 @@ func TestEnvironmentWorkerRunHandlesPolledWorkInProcess(t *testing.T) { if api.ackCount != 1 || api.stopCount != 2 { t.Fatalf("ack_count=%d stop_count=%d", api.ackCount, api.stopCount) } - if !api.stops[0].Force { + if force, ok := api.stops[0].Force.Get(); !ok || !force { t.Fatalf("worker stop should be force=true: %+v", api.stops[0]) } - if api.stops[1].Force { + if _, ok := api.stops[1].Force.Get(); ok { t.Fatalf("poller release stop should be force=false: %+v", api.stops[1]) } if len(api.sent) != 1 { @@ -166,12 +161,7 @@ func TestEnvironmentWorkerStopsWorkOnSessionIdleEvent(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() api := &fakeEnvironmentWorkerAPI{ - pollItem: &selfhosted.WorkItem{ - ID: "work_idle", - EnvironmentID: "env_local", - SessionID: "sess_idle", - LeaseID: "lease_idle", - }, + pollItem: newTestWorkItem("work_idle", "env_local", "sess_idle"), events: []selfhosted.Event{{ ID: "evt_idle", Type: selfhosted.EventTypeSessionStatusIdle, @@ -193,8 +183,10 @@ func TestEnvironmentWorkerStopsWorkOnSessionIdleEvent(t *testing.T) { if api.ackCount != 1 || api.stopCount == 0 { t.Fatalf("ack_count=%d stop_count=%d", api.ackCount, api.stopCount) } - if got := api.stops[0]; !got.Force || got.WorkID != "work_idle" { + if got := api.stops[0]; got.WorkID != "work_idle" { t.Fatalf("worker stop = %+v", got) + } else if force, ok := got.Force.Get(); !ok || !force { + t.Fatalf("worker stop should be force=true: %+v", got) } if len(api.sent) != 0 { t.Fatalf("sent events = %+v", api.sent) @@ -244,22 +236,17 @@ func TestEnvironmentWorkerHeartbeatLeaseLostOnPreconditionFailed(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() var cause heartbeatStopCause - w.heartbeatLoop(ctx, selfhosted.WorkItem{ - ID: "work_local", - EnvironmentID: "env_local", - SessionID: "sess_local", - TTLSeconds: 1, - }, api, cancel, func(c heartbeatStopCause) { + w.heartbeatLoop(ctx, claimedWork{ID: "work_local", EnvironmentID: "env_local", SessionID: "sess_local"}, api, cancel, func(c heartbeatStopCause) { cause = c }) if cause != heartbeatStopCauseLeaseLost { t.Fatalf("heartbeat cause = %q", cause) } - if got.ExpectedLastHeartbeat != selfhosted.ExpectedLastHeartbeatNoHeartbeat { - t.Fatalf("expected_last_heartbeat=%q", got.ExpectedLastHeartbeat) + if value, ok := got.ExpectedLastHeartbeat.Get(); !ok || value != selfhosted.ExpectedLastHeartbeatNoHeartbeat { + t.Fatalf("expected_last_heartbeat=%+v", got.ExpectedLastHeartbeat) } - if got.DesiredTTLSeconds != int(heartbeatDefault/time.Second) { - t.Fatalf("desired_ttl_seconds=%d", got.DesiredTTLSeconds) + if value, ok := got.DesiredTTLSeconds.Get(); !ok || value != int64(heartbeatDefault/time.Second) { + t.Fatalf("desired_ttl_seconds=%+v", got.DesiredTTLSeconds) } } @@ -357,11 +344,7 @@ func TestEnvironmentWorkerHeartbeatStopsOnPermanent4xx(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() var cause heartbeatStopCause - w.heartbeatLoop(ctx, selfhosted.WorkItem{ - ID: "work_local", - EnvironmentID: "env_local", - SessionID: "sess_local", - }, api, cancel, func(c heartbeatStopCause) { + w.heartbeatLoop(ctx, claimedWork{ID: "work_local", EnvironmentID: "env_local", SessionID: "sess_local"}, api, cancel, func(c heartbeatStopCause) { cause = c }) if cause != heartbeatStopCausePermanentFailure { @@ -377,7 +360,7 @@ func TestEnvironmentWorkerHeartbeatPrefersStopRequestedStateOverLeaseNotExtended heartbeat: func(context.Context, selfhosted.HeartbeatWorkRequest) (*selfhosted.HeartbeatResponse, error) { return &selfhosted.HeartbeatResponse{ LastHeartbeat: time.Now().UTC().Format(time.RFC3339Nano), - LeaseExtended: selfhosted.BoolPtr(false), + LeaseExtended: false, State: selfhosted.WorkStateStopping, }, nil }, @@ -390,11 +373,7 @@ func TestEnvironmentWorkerHeartbeatPrefersStopRequestedStateOverLeaseNotExtended ctx, cancel := context.WithCancel(context.Background()) defer cancel() var cause heartbeatStopCause - w.heartbeatLoop(ctx, selfhosted.WorkItem{ - ID: "work_local", - EnvironmentID: "env_local", - SessionID: "sess_local", - }, api, cancel, func(c heartbeatStopCause) { + w.heartbeatLoop(ctx, claimedWork{ID: "work_local", EnvironmentID: "env_local", SessionID: "sess_local"}, api, cancel, func(c heartbeatStopCause) { cause = c }) if cause != heartbeatStopCauseStopRequested { @@ -426,13 +405,9 @@ func TestEnvironmentWorkerHeartbeatUsesAnthropicDefaultTTL(t *testing.T) { }) ctx, cancel := context.WithCancel(context.Background()) defer cancel() - w.heartbeatLoop(ctx, selfhosted.WorkItem{ - ID: "work_local", - EnvironmentID: "env_local", - SessionID: "sess_local", - }, api, cancel, nil) - if got.DesiredTTLSeconds != int(heartbeatDefault/time.Second) { - t.Fatalf("desired_ttl_seconds=%d", got.DesiredTTLSeconds) + w.heartbeatLoop(ctx, claimedWork{ID: "work_local", EnvironmentID: "env_local", SessionID: "sess_local"}, api, cancel, nil) + if value, ok := got.DesiredTTLSeconds.Get(); !ok || value != int64(heartbeatDefault/time.Second) { + t.Fatalf("desired_ttl_seconds=%+v", got.DesiredTTLSeconds) } if requestTimeout < heartbeatDefault/2-time.Second || requestTimeout > heartbeatDefault/2 { t.Fatalf("heartbeat request timeout=%s", requestTimeout) diff --git a/arkruntime/lib/environments/poller.go b/arkruntime/lib/environments/poller.go index c7c8380..74a003b 100644 --- a/arkruntime/lib/environments/poller.go +++ b/arkruntime/lib/environments/poller.go @@ -17,6 +17,7 @@ import ( "github.com/volcengine/ark-runtime-go/arkruntime" "github.com/volcengine/ark-runtime-go/arkruntime/internal/selfhostedlog" + "github.com/volcengine/ark-runtime-go/arkruntime/model/environment" selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" ) @@ -131,7 +132,7 @@ func (p *WorkPoller) Next() bool { if err := p.api.AckWork(p.ctx, selfhosted.AckWorkRequest{ EnvironmentID: item.EnvironmentID, WorkID: item.ID, - WorkerID: p.opts.WorkerID, + WorkerID: environment.NewOptString(p.opts.WorkerID), }); err != nil { p.logger.Warn("ack work failed", "work_id", item.ID, "err", err) // ACK 是 queued -> starting 的竞争操作。失败时无法证明当前 @@ -204,7 +205,7 @@ func (p *WorkPoller) discardInvalidWork(item selfhosted.WorkItem, _ string) { if err := p.api.AckWork(p.ctx, selfhosted.AckWorkRequest{ EnvironmentID: item.EnvironmentID, WorkID: item.ID, - WorkerID: p.opts.WorkerID, + WorkerID: environment.NewOptString(p.opts.WorkerID), }); err != nil { p.logger.Warn("ack invalid work failed", "work_id", item.ID, "err", err) return @@ -214,7 +215,7 @@ func (p *WorkPoller) discardInvalidWork(item selfhosted.WorkItem, _ string) { if err := p.api.StopWork(ctx, selfhosted.StopWorkRequest{ EnvironmentID: item.EnvironmentID, WorkID: item.ID, - Force: true, + Force: environment.NewOptBool(true), }); err != nil && !isResolvedStatus(err) { p.logger.Warn("stop invalid work failed", "work_id", item.ID, "err", err) } diff --git a/arkruntime/lib/environments/poller_test.go b/arkruntime/lib/environments/poller_test.go index 6792e5f..2e6627a 100644 --- a/arkruntime/lib/environments/poller_test.go +++ b/arkruntime/lib/environments/poller_test.go @@ -8,6 +8,7 @@ import ( "net/http" "testing" + "github.com/volcengine/ark-runtime-go/arkruntime/model/environment" selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" ) @@ -24,6 +25,20 @@ type fakePollerAPI struct { stops []selfhosted.StopWorkRequest } +func newTestWorkItem(workID, environmentID, sessionID string) *selfhosted.WorkItem { + return &selfhosted.WorkItem{ + ID: workID, + CreatedAt: "2026-08-24T10:00:00Z", + EnvironmentID: environmentID, + Data: selfhosted.WorkData{ + ID: sessionID, + Type: "session", + }, + State: environment.WorkStateActive, + Type: environment.WorkItemTypeWork, + } +} + func (f *fakePollerAPI) PollWork(context.Context, selfhosted.PollWorkRequest) (*selfhosted.WorkItem, error) { item := f.pollItem f.pollItem = nil @@ -63,11 +78,7 @@ func (f *fakePollerAPI) OpenSkill(context.Context, selfhosted.OpenSkillRequest) func TestWorkPollerAcksAndStopsOnClose(t *testing.T) { api := &fakePollerAPI{ - pollItem: &selfhosted.WorkItem{ - ID: testWorkID, - EnvironmentID: "env_1", - SessionID: testSessionID, - }, + pollItem: newTestWorkItem(testWorkID, "env_1", testSessionID), } poller := NewWorkPoller(context.Background(), api, WorkPollerOptions{ EnvironmentID: "env_1", @@ -92,18 +103,14 @@ func TestWorkPollerAcksAndStopsOnClose(t *testing.T) { if api.stops[0].WorkID != testWorkID { t.Fatalf("stop request = %+v", api.stops[0]) } - if api.stops[0].Force { + if _, ok := api.stops[0].Force.Get(); ok { t.Fatalf("poller release should not force stop: %+v", api.stops[0]) } } func TestWorkPollerStopsPreviousBeforeNext(t *testing.T) { api := &fakePollerAPI{ - pollItem: &selfhosted.WorkItem{ - ID: testWorkID, - EnvironmentID: "env_1", - SessionID: testSessionID, - }, + pollItem: newTestWorkItem(testWorkID, "env_1", testSessionID), } poller := NewWorkPoller(context.Background(), api, WorkPollerOptions{ EnvironmentID: "env_1", @@ -113,11 +120,7 @@ func TestWorkPollerStopsPreviousBeforeNext(t *testing.T) { if !poller.Next() { t.Fatalf("Next returned false: %v", poller.Err()) } - api.pollItem = &selfhosted.WorkItem{ - ID: "work_2", - EnvironmentID: "env_1", - SessionID: "sess_2", - } + api.pollItem = newTestWorkItem("work_2", "env_1", "sess_2") if !poller.Next() { t.Fatalf("Next returned false: %v", poller.Err()) } @@ -134,12 +137,8 @@ func TestWorkPollerStopsPreviousBeforeNext(t *testing.T) { func TestWorkPollerDoesNotStopWhenAckLosesClaimRace(t *testing.T) { api := &fakePollerAPI{ - pollItem: &selfhosted.WorkItem{ - ID: testWorkID, - EnvironmentID: "env_1", - SessionID: testSessionID, - }, - ackErr: &selfhosted.APIError{StatusCode: 409, Message: "already claimed"}, + pollItem: newTestWorkItem(testWorkID, "env_1", testSessionID), + ackErr: &selfhosted.APIError{StatusCode: 409, Message: "already claimed"}, } poller := NewWorkPoller(context.Background(), api, WorkPollerOptions{ EnvironmentID: "env_1", @@ -165,12 +164,8 @@ func TestPollerConflictIsRecoverable(t *testing.T) { func TestWorkPollerStopsPollingOnPermanentAckFailure(t *testing.T) { api := &fakePollerAPI{ - pollItem: &selfhosted.WorkItem{ - ID: testWorkID, - EnvironmentID: "env_1", - SessionID: testSessionID, - }, - ackErr: &selfhosted.APIError{StatusCode: 401, Message: "invalid credential"}, + pollItem: newTestWorkItem(testWorkID, "env_1", testSessionID), + ackErr: &selfhosted.APIError{StatusCode: 401, Message: "invalid credential"}, } poller := NewWorkPoller(context.Background(), api, WorkPollerOptions{ EnvironmentID: "env_1", @@ -190,10 +185,7 @@ func TestWorkPollerStopsPollingOnPermanentAckFailure(t *testing.T) { func TestWorkPollerTreatsEmptyWorkIDAsEmptyPoll(t *testing.T) { api := &fakePollerAPI{ - pollItem: &selfhosted.WorkItem{ - EnvironmentID: "env_1", - SessionID: testSessionID, - }, + pollItem: newTestWorkItem("", "env_1", testSessionID), } poller := NewWorkPoller(context.Background(), api, WorkPollerOptions{ EnvironmentID: "env_1", diff --git a/arkruntime/lib/environments/worker.go b/arkruntime/lib/environments/worker.go index 47af04c..e4b91ff 100644 --- a/arkruntime/lib/environments/worker.go +++ b/arkruntime/lib/environments/worker.go @@ -15,6 +15,7 @@ import ( "github.com/volcengine/ark-runtime-go/arkruntime" "github.com/volcengine/ark-runtime-go/arkruntime/internal/selfhostedlog" + "github.com/volcengine/ark-runtime-go/arkruntime/model/environment" selfhosted "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted/envinit" "github.com/volcengine/ark-runtime-go/arkruntime/tools/agenttoolset" @@ -106,9 +107,14 @@ func (w *EnvironmentWorker) Run(ctx context.Context) error { } func (w *EnvironmentWorker) runClaimedWork(ctx context.Context, item selfhosted.WorkItem, logger *selfhostedlog.Logger) { - if err := w.handleItem(ctx, item, false); err != nil && + work, err := claimedWorkFromItem(item, w.opts.EnvironmentID) + if err != nil { + logger.Warn("handle work failed", "work_id", item.ID, "err", err) + return + } + if err := w.handleItem(ctx, work, false); err != nil && !isBenignWorkerExit(err) { - logger.Warn("handle work failed", "work_id", item.ID, "session_id", item.SessionIDValue(), "err", err) + logger.Warn("handle work failed", "work_id", work.ID, "session_id", work.SessionID, "err", err) } } @@ -125,37 +131,27 @@ func (w *EnvironmentWorker) HandleItem(ctx context.Context, opts HandleItemOptio if ctx == nil { ctx = context.Background() } - item, err := w.resolveHandleItem(opts) + work, err := claimedWorkFromOptions(opts) if err != nil { return err } - if err := w.handleItem(ctx, item, true); err != nil && + if err := w.handleItem(ctx, work, true); err != nil && !isBenignWorkerExit(err) { return err } return nil } -func (w *EnvironmentWorker) handleItem(ctx context.Context, item selfhosted.WorkItem, useWorkdirAsSession bool) (err error) { +func (w *EnvironmentWorker) handleItem(ctx context.Context, work claimedWork, useWorkdirAsSession bool) (err error) { if w.api == nil { return errors.New("environments: API is required") } - if item.EnvironmentID == "" { - item.EnvironmentID = firstNonEmpty(w.opts.EnvironmentID, os.Getenv("MA_ENVIRONMENT_ID")) - } - if item.ID == "" { - return errors.New("work item id must not be empty") - } - sessionID := item.SessionIDValue() - if sessionID == "" { - return errors.New("work item does not contain session id") - } - workdir, err := w.workdirFor(sessionID, useWorkdirAsSession) + workdir, err := w.workdirFor(work.SessionID, useWorkdirAsSession) if err != nil { return err } api := w.api - logger := w.logger().With("component", "environment-worker", "work_id", item.ID, "session_id", sessionID, "workdir", workdir) + logger := w.logger().With("component", "environment-worker", "work_id", work.ID, "session_id", work.SessionID, "workdir", workdir) workCtx, cancel := context.WithCancel(ctx) defer cancel() @@ -163,7 +159,7 @@ func (w *EnvironmentWorker) handleItem(ctx context.Context, item selfhosted.Work heartbeatDone := make(chan struct{}) go func() { defer close(heartbeatDone) - w.heartbeatLoop(workCtx, item, api, cancel, func(cause heartbeatStopCause) { + w.heartbeatLoop(workCtx, work, api, cancel, func(cause heartbeatStopCause) { heartbeatCause.Store(string(cause)) }) }() @@ -173,7 +169,7 @@ func (w *EnvironmentWorker) handleItem(ctx context.Context, item selfhosted.Work <-heartbeatDone cause := loadHeartbeatStopCause(&heartbeatCause) if shouldStopItem(cause) { - _ = w.stopItem(api, item) + _ = w.stopItem(api, work) } else { logger.Info("skip stop work after heartbeat ownership became uncertain", "cause", cause) } @@ -182,7 +178,7 @@ func (w *EnvironmentWorker) handleItem(ctx context.Context, item selfhosted.Work } }() - session, err := api.GetSession(workCtx, selfhosted.GetSessionRequest{SessionID: sessionID}) + session, err := api.GetSession(workCtx, selfhosted.GetSessionRequest{SessionID: work.SessionID}) if err != nil { return fmt.Errorf("get session: %w", err) } @@ -191,7 +187,7 @@ func (w *EnvironmentWorker) handleItem(ctx context.Context, item selfhosted.Work return err } if session.ID == "" { - session.ID = sessionID + session.ID = work.SessionID } toolEnv := w.toolContext(workdir) initOpts := toolEnv.InitOptions() @@ -211,8 +207,8 @@ func (w *EnvironmentWorker) handleItem(ctx context.Context, item selfhosted.Work if err != nil { return fmt.Errorf("create tool result store: %w", err) } - runner := selfhosted.NewSessionToolRunner(workCtx, api, sessionID, selfhosted.SessionToolRunnerOptions{ - WorkID: item.ID, + runner := selfhosted.NewSessionToolRunner(workCtx, api, work.SessionID, selfhosted.SessionToolRunnerOptions{ + WorkID: work.ID, Tools: tools, CustomTools: w.opts.CustomTools, ResultStore: store, @@ -289,28 +285,20 @@ func (w *EnvironmentWorker) workdirFor(sessionID string, useWorkdirAsSession boo return sessionWorkdir(absRoot, sessionID) } -func (w *EnvironmentWorker) resolveHandleItem(opts HandleItemOptions) (selfhosted.WorkItem, error) { - item, err := workItemFromOptions(opts) - if err != nil { - return item, err - } - return item, nil -} - -func (w *EnvironmentWorker) stopItem(api selfhosted.API, item selfhosted.WorkItem) error { +func (w *EnvironmentWorker) stopItem(api selfhosted.API, work claimedWork) error { stopCtx, stopCancel := context.WithTimeout(context.Background(), stopTimeout) defer stopCancel() req := selfhosted.StopWorkRequest{ - EnvironmentID: item.EnvironmentID, - WorkID: item.ID, - Force: true, + EnvironmentID: work.EnvironmentID, + WorkID: work.ID, + Force: environment.NewOptBool(true), } if err := api.StopWork(stopCtx, req); err != nil { if selfhosted.IsStatus(err, 409) || selfhosted.IsStatus(err, 412) { - w.logger().Info("stop work already resolved", "work_id", item.ID, "err", err) + w.logger().Info("stop work already resolved", "work_id", work.ID, "err", err) return nil } - w.logger().Warn("stop work failed", "work_id", item.ID, "err", err) + w.logger().Warn("stop work failed", "work_id", work.ID, "err", err) return err } return nil @@ -348,29 +336,47 @@ func isBenignWorkerExit(err error) bool { errors.Is(err, selfhosted.ErrSessionTerminated) } -func workItemFromOptions(opts HandleItemOptions) (selfhosted.WorkItem, error) { +type claimedWork struct { + ID string + EnvironmentID string + SessionID string + LatestHeartbeatAt string +} + +func claimedWorkFromItem(item selfhosted.WorkItem, fallbackEnvironmentID string) (claimedWork, error) { + work := claimedWork{ + ID: item.ID, + EnvironmentID: firstNonEmpty(item.EnvironmentID, fallbackEnvironmentID), + SessionID: item.SessionIDValue(), + LatestHeartbeatAt: item.LatestHeartbeatValue(), + } + return validateClaimedWork(work) +} + +func claimedWorkFromOptions(opts HandleItemOptions) (claimedWork, error) { workID := firstNonEmpty(opts.WorkID, os.Getenv("MA_WORK_ID")) environmentID := firstNonEmpty(opts.EnvironmentID, os.Getenv("MA_ENVIRONMENT_ID")) sessionID := firstNonEmpty(opts.SessionID, os.Getenv("MA_SESSION_ID")) latestHeartbeatAt := firstNonEmpty(opts.LatestHeartbeatAt, os.Getenv("MA_LATEST_HEARTBEAT_AT")) - if workID == "" { - return selfhosted.WorkItem{}, errors.New("environments: work id is required") - } - if environmentID == "" { - return selfhosted.WorkItem{}, errors.New("environments: environment id is required") - } - if sessionID == "" { - return selfhosted.WorkItem{}, errors.New("environments: session id is required") - } - return selfhosted.WorkItem{ + return validateClaimedWork(claimedWork{ ID: workID, EnvironmentID: environmentID, + SessionID: sessionID, LatestHeartbeatAt: latestHeartbeatAt, - Data: selfhosted.WorkData{ - Type: "session", - ID: sessionID, - }, - }, nil + }) +} + +func validateClaimedWork(work claimedWork) (claimedWork, error) { + if work.ID == "" { + return claimedWork{}, errors.New("environments: work id is required") + } + if work.EnvironmentID == "" { + return claimedWork{}, errors.New("environments: environment id is required") + } + if work.SessionID == "" { + return claimedWork{}, errors.New("environments: session id is required") + } + return work, nil } func firstNonEmpty(values ...string) string { diff --git a/arkruntime/lib/environments/worker_test.go b/arkruntime/lib/environments/worker_test.go index c03bc4e..adfcee4 100644 --- a/arkruntime/lib/environments/worker_test.go +++ b/arkruntime/lib/environments/worker_test.go @@ -4,8 +4,8 @@ package environments import "testing" -func TestHandleItemOptionsBuildsWorkItem(t *testing.T) { - got, err := workItemFromOptions(HandleItemOptions{ +func TestHandleItemOptionsBuildsClaimedWork(t *testing.T) { + got, err := claimedWorkFromOptions(HandleItemOptions{ WorkID: testWorkID, EnvironmentID: "env_1", SessionID: testSessionID, @@ -14,15 +14,12 @@ func TestHandleItemOptionsBuildsWorkItem(t *testing.T) { if err != nil { t.Fatal(err) } - if got.ID != testWorkID || got.EnvironmentID != "env_1" || got.SessionIDValue() != testSessionID { - t.Fatalf("work item = %+v", got) + if got.ID != testWorkID || got.EnvironmentID != "env_1" || got.SessionID != testSessionID { + t.Fatalf("claimed work = %+v", got) } - if got.LatestHeartbeatValue() != "2026-08-11T00:00:00Z" { + if got.LatestHeartbeatAt != "2026-08-11T00:00:00Z" { t.Fatalf("work heartbeat = %+v", got) } - if got.Data.Type != "session" || got.Data.ID != testSessionID { - t.Fatalf("work data = %+v", got.Data) - } } func TestHandleItemOptionsUsesEnvironmentFallbacks(t *testing.T) { @@ -31,20 +28,20 @@ func TestHandleItemOptionsUsesEnvironmentFallbacks(t *testing.T) { t.Setenv("MA_SESSION_ID", "sess_env") t.Setenv("MA_LATEST_HEARTBEAT_AT", "2026-08-11T00:00:00Z") - got, err := workItemFromOptions(HandleItemOptions{}) + got, err := claimedWorkFromOptions(HandleItemOptions{}) if err != nil { t.Fatal(err) } - if got.ID != "work_env" || got.EnvironmentID != "env_env" || got.SessionIDValue() != "sess_env" { - t.Fatalf("work item = %+v", got) + if got.ID != "work_env" || got.EnvironmentID != "env_env" || got.SessionID != "sess_env" { + t.Fatalf("claimed work = %+v", got) } - if got.LatestHeartbeatValue() != "2026-08-11T00:00:00Z" { + if got.LatestHeartbeatAt != "2026-08-11T00:00:00Z" { t.Fatalf("work heartbeat = %+v", got) } } func TestHandleItemOptionsRequiresWorkID(t *testing.T) { - if _, err := workItemFromOptions(HandleItemOptions{EnvironmentID: "env_1", SessionID: "sess_1"}); err == nil { + if _, err := claimedWorkFromOptions(HandleItemOptions{EnvironmentID: "env_1", SessionID: "sess_1"}); err == nil { t.Fatal("expected work id error") } } diff --git a/arkruntime/selfhosted/client_api.go b/arkruntime/selfhosted/client_api.go index 99aa054..d08b36e 100644 --- a/arkruntime/selfhosted/client_api.go +++ b/arkruntime/selfhosted/client_api.go @@ -35,19 +35,11 @@ func (a *ClientAPI) PollWork(ctx context.Context, req PollWorkRequest) (*WorkIte if a == nil || a.client == nil { return nil, errors.New("selfhosted: arkruntime client is nil") } - item, err := a.client.PollWork(ctx, &environment.PollWorkRequest{ - EnvironmentID: req.EnvironmentID, - WorkerID: req.WorkerID, - MaxItems: req.MaxItems, - BlockMS: req.BlockMS, - ReclaimOlderThanMS: req.ReclaimOlderThanMS, - WorkerClientType: firstNonEmpty(req.WorkerClientType, DefaultWorkerClientType), - WorkerClientVersion: firstNonEmpty(req.WorkerClientVersion, DefaultWorkerClientVersion), - }) + item, err := a.client.PollWork(ctx, &req) if err != nil { return nil, toWorkerAPIError(err) } - return fromEnvironmentWorkItem(item), nil + return item, nil } // AckWork acknowledges one claimed work item. @@ -55,11 +47,7 @@ func (a *ClientAPI) AckWork(ctx context.Context, req AckWorkRequest) error { if a == nil || a.client == nil { return errors.New("selfhosted: arkruntime client is nil") } - err := a.client.AckWork(ctx, &environment.AckWorkRequest{ - EnvironmentID: req.EnvironmentID, - WorkID: req.WorkID, - WorkerID: envOptString(req.WorkerID), - }) + err := a.client.AckWork(ctx, &req) return toWorkerAPIError(err) } @@ -68,29 +56,15 @@ func (a *ClientAPI) HeartbeatWork(ctx context.Context, req HeartbeatWorkRequest) if a == nil || a.client == nil { return nil, errors.New("selfhosted: arkruntime client is nil") } - expectedLastHeartbeat := req.ExpectedLastHeartbeat - if expectedLastHeartbeat == "" { - expectedLastHeartbeat = ExpectedLastHeartbeatNoHeartbeat + expectedLastHeartbeat, ok := req.ExpectedLastHeartbeat.Get() + if !ok || expectedLastHeartbeat == "" { + req.ExpectedLastHeartbeat = environment.NewOptString(ExpectedLastHeartbeatNoHeartbeat) } - resp, err := a.client.HeartbeatWork(ctx, &environment.HeartbeatWorkRequest{ - EnvironmentID: req.EnvironmentID, - WorkID: req.WorkID, - ExpectedLastHeartbeat: envOptString(expectedLastHeartbeat), - DesiredTTLSeconds: envOptInt64(req.DesiredTTLSeconds), - }) + resp, err := a.client.HeartbeatWork(ctx, &req) if err != nil { return nil, toWorkerAPIError(err) } - if resp == nil { - return nil, nil - } - return &HeartbeatResponse{ - LastHeartbeat: resp.LastHeartbeat, - LeaseExtended: optionalBoolPtr(resp.LeaseExtended, true), - State: string(resp.State), - TTLSeconds: int(resp.TTLSeconds), - Type: string(resp.Type), - }, nil + return resp, nil } // StopWork releases or stops one claimed work item. @@ -98,11 +72,7 @@ func (a *ClientAPI) StopWork(ctx context.Context, req StopWorkRequest) error { if a == nil || a.client == nil { return errors.New("selfhosted: arkruntime client is nil") } - err := a.client.StopWork(ctx, &environment.StopWorkRequest{ - EnvironmentID: req.EnvironmentID, - WorkID: req.WorkID, - Force: envOptBool(req.Force), - }) + err := a.client.StopWork(ctx, &req) return toWorkerAPIError(err) } @@ -290,50 +260,6 @@ func (a *ClientAPI) OpenSkill(ctx context.Context, req OpenSkillRequest) (*Skill }, nil } -func fromEnvironmentWorkItem(item *environment.WorkItem) *WorkItem { - if item == nil { - return nil - } - latestHeartbeat := item.LatestHeartbeatValue() - acknowledgedAt, _ := item.AcknowledgedAt.Get() - secret, _ := item.Secret.Get() - startedAt, _ := item.StartedAt.Get() - stopRequestedAt, _ := item.StopRequestedAt.Get() - stoppedAt, _ := item.StoppedAt.Get() - data := WorkData{} - if item.Data.ID != "" || item.Data.Type != "" { - data = WorkData{ - Type: item.Data.Type, - ID: item.Data.ID, - } - } - tags := make([]WorkTag, 0, len(item.Tags)) - for _, tag := range item.Tags { - value, _ := tag.Value.Get() - tags = append(tags, WorkTag{ - Key: tag.Key, - Value: value, - }) - } - return &WorkItem{ - ID: item.ID, - AcknowledgedAt: acknowledgedAt, - CreatedAt: item.CreatedAt, - Data: data, - EnvironmentID: item.EnvironmentID, - LatestHeartbeatAt: latestHeartbeat, - Tags: tags, - Secret: secret, - StartedAt: startedAt, - State: string(item.State), - StopRequestedAt: stopRequestedAt, - StoppedAt: stoppedAt, - Type: string(item.Type), - SessionID: data.ID, - LastHeartbeat: latestHeartbeat, - } -} - func fromRuntimeSession(in *session.Session) *Session { if in == nil { return nil @@ -449,33 +375,5 @@ func firstNonEmpty(values ...string) string { return "" } -func envOptString(value string) environment.OptString { - if value == "" { - return environment.OptString{} - } - return environment.NewOptString(value) -} - -func envOptInt64(value int) environment.OptInt64 { - if value <= 0 { - return environment.OptInt64{} - } - return environment.NewOptInt64(int64(value)) -} - -func envOptBool(value bool) environment.OptBool { - if !value { - return environment.OptBool{} - } - return environment.NewOptBool(value) -} - -func optionalBoolPtr(value bool, ok bool) *bool { - if !ok { - return nil - } - return &value -} - var _ API = (*ClientAPI)(nil) var _ EventStreamer = (*ClientAPI)(nil) diff --git a/arkruntime/selfhosted/client_api_test.go b/arkruntime/selfhosted/client_api_test.go index 447b179..dcee2de 100644 --- a/arkruntime/selfhosted/client_api_test.go +++ b/arkruntime/selfhosted/client_api_test.go @@ -12,6 +12,7 @@ import ( "testing" "github.com/volcengine/ark-runtime-go/arkruntime" + "github.com/volcengine/ark-runtime-go/arkruntime/model/environment" "github.com/volcengine/ark-runtime-go/arkruntime/selfhosted" ) @@ -181,7 +182,7 @@ func TestClientAPIHeartbeatDefaultsExpectedLastHeartbeat(t *testing.T) { resp, err := api.HeartbeatWork(context.Background(), selfhosted.HeartbeatWorkRequest{ EnvironmentID: "env-1", WorkID: "work-1", - DesiredTTLSeconds: 30, + DesiredTTLSeconds: environment.NewOptInt64(30), }) if err != nil { t.Fatalf("HeartbeatWork() error = %v", err) diff --git a/arkruntime/selfhosted/types.go b/arkruntime/selfhosted/types.go index 89692b2..95f60d5 100644 --- a/arkruntime/selfhosted/types.go +++ b/arkruntime/selfhosted/types.go @@ -12,6 +12,8 @@ import ( "fmt" "io" "time" + + "github.com/volcengine/ark-runtime-go/arkruntime/model/environment" ) const ( @@ -67,19 +69,15 @@ const ( const ( // WorkStateQueued 表示 work 仍在队列中。 - WorkStateQueued = "queued" + WorkStateQueued = environment.WorkStateQueued // WorkStateStarting 表示 work 已 ack 并正在启动执行环境。 - WorkStateStarting = "starting" + WorkStateStarting = environment.WorkStateStarting // WorkStateActive 表示 work lease 仍归当前 worker 所有。 - WorkStateActive = "active" - // WorkStateRunning 兼容 MA 早期返回的 running 状态。 - // - // Deprecated: MA work state 不再包含 running,使用 active。 - WorkStateRunning = "running" + WorkStateActive = environment.WorkStateActive // WorkStateStopping 表示控制面要求 worker 停止当前 work。 - WorkStateStopping = "stopping" + WorkStateStopping = environment.WorkStateStopping // WorkStateStopped 表示当前 work 已经停止。 - WorkStateStopped = "stopped" + WorkStateStopped = environment.WorkStateStopped ) const ( @@ -106,17 +104,6 @@ const ( WorkerErrorKindInvalidResponse = "invalid_response" ) -const ( - // DefaultWorkerClientType 保留用于源码兼容。 - // - // Deprecated: MA work API 不再接收 worker client header。 - DefaultWorkerClientType = "ark-self-hosted-worker-go" - // DefaultWorkerClientVersion 保留用于源码兼容。 - // - // Deprecated: MA work API 不再接收 worker client header。 - DefaultWorkerClientVersion = "0.1.0" -) - // ErrEventStreamUnsupported 表示当前 API 实现不支持 SSE 事件流。 var ErrEventStreamUnsupported = errors.New("ark event stream is not configured") @@ -138,77 +125,19 @@ type SkillResolver interface { } // PollWorkRequest 是 worker 轮询 work queue 的请求。 -type PollWorkRequest struct { - EnvironmentID string `json:"environment_id"` - WorkerID string `json:"worker_id,omitempty"` - // Deprecated: MA poll 每次最多返回一个 work。 - MaxItems int `json:"max_items,omitempty"` - BlockMS int `json:"block_ms,omitempty"` - ReclaimOlderThanMS int `json:"reclaim_older_than_ms,omitempty"` - // Deprecated: MA work API 不再接收 worker client header。 - WorkerClientType string `json:"-"` - // Deprecated: MA work API 不再接收 worker client header。 - WorkerClientVersion string `json:"-"` -} +type PollWorkRequest = environment.PollWorkRequest // AckWorkRequest 是 worker 确认已接收 work 的请求。 -type AckWorkRequest struct { - EnvironmentID string `json:"environment_id"` - WorkID string `json:"work_id"` - WorkerID string `json:"worker_id,omitempty"` - // Deprecated: MA ack 接口不再接收 lease_id。 - LeaseID string `json:"lease_id,omitempty"` -} +type AckWorkRequest = environment.AckWorkRequest // HeartbeatWorkRequest 是 worker 刷新 work lease 的请求。 -type HeartbeatWorkRequest struct { - EnvironmentID string `json:"environment_id"` - WorkID string `json:"work_id"` - ExpectedLastHeartbeat string `json:"expected_last_heartbeat,omitempty"` - DesiredTTLSeconds int `json:"desired_ttl_seconds,omitempty"` - // Deprecated: MA heartbeat 接口不再接收 lease_id。 - LeaseID string `json:"lease_id,omitempty"` - // Deprecated: MA heartbeat 接口不再接收 worker_id。 - WorkerID string `json:"worker_id,omitempty"` - // Deprecated: 使用 ExpectedLastHeartbeat 表达 CAS 期望。 - LastHeartbeat string `json:"last_heartbeat,omitempty"` - // Deprecated: MA heartbeat 接口不再接收 active_tool_call_id。 - ActiveToolCallID string `json:"active_tool_call_id,omitempty"` -} +type HeartbeatWorkRequest = environment.HeartbeatWorkRequest // HeartbeatResponse 是控制面返回的 heartbeat 结果。 -type HeartbeatResponse struct { - LastHeartbeat string `json:"last_heartbeat,omitempty"` - LeaseExtended *bool `json:"lease_extended,omitempty"` - State string `json:"state,omitempty"` - TTLSeconds int `json:"ttl_seconds,omitempty"` - Type string `json:"type,omitempty"` - // Deprecated: MA heartbeat 响应不再返回 lease_expires_at。 - LeaseExpiresAt string `json:"lease_expires_at,omitempty"` - // Deprecated: 使用 TTLSeconds。 - LeaseSeconds int `json:"lease_seconds,omitempty"` - // Deprecated: 使用 State。 - Status string `json:"status,omitempty"` - // Deprecated: 使用 State 判断 stopping/stopped。 - StopRequested bool `json:"stop_requested,omitempty"` - // Deprecated: MA heartbeat 响应体不再返回 request_id。 - RequestID string `json:"request_id,omitempty"` -} +type HeartbeatResponse = environment.HeartbeatWorkResponse // StopWorkRequest 是 worker 请求控制面结束或停止 work 的请求。 -type StopWorkRequest struct { - EnvironmentID string `json:"environment_id"` - WorkID string `json:"work_id"` - Force bool `json:"force,omitempty"` - // Deprecated: MA stop 接口不再接收 lease_id。 - LeaseID string `json:"lease_id,omitempty"` - // Deprecated: MA stop 接口不再接收 worker_id。 - WorkerID string `json:"worker_id,omitempty"` - // Deprecated: MA stop 接口只接收 force。 - Reason string `json:"reason,omitempty"` - // Deprecated: MA stop 接口只接收 force。 - Message string `json:"message,omitempty"` -} +type StopWorkRequest = environment.StopWorkRequest // GetSessionRequest 是读取 session 配置的请求。 type GetSessionRequest struct { @@ -254,83 +183,13 @@ type OpenSkillRequest struct { } // WorkItem 表示控制面分配给 worker 的一份工作。 -type WorkItem struct { - ID string `json:"id"` - AcknowledgedAt string `json:"acknowledged_at,omitempty"` - CreatedAt string `json:"created_at,omitempty"` - Data WorkData `json:"data,omitempty"` - EnvironmentID string `json:"environment_id,omitempty"` - LatestHeartbeatAt string `json:"latest_heartbeat_at,omitempty"` - Tags []WorkTag `json:"tags,omitempty"` - Secret string `json:"secret,omitempty"` - StartedAt string `json:"started_at,omitempty"` - State string `json:"state,omitempty"` - StopRequestedAt string `json:"stop_requested_at,omitempty"` - StoppedAt string `json:"stopped_at,omitempty"` - Type string `json:"type,omitempty"` - // Deprecated: 使用 Data.ID。 - SessionID string `json:"session_id,omitempty"` - // Deprecated: 使用 State。 - Status string `json:"status,omitempty"` - // Deprecated: MA work item 不再返回 attempt。 - Attempt int `json:"attempt,omitempty"` - // Deprecated: MA work item 不再返回 lease_id。 - LeaseID string `json:"lease_id,omitempty"` - // Deprecated: MA work item 不再返回 lease_expires_at。 - LeaseExpiresAt string `json:"lease_expires_at,omitempty"` - // Deprecated: MA work item 不再返回 request_id。 - RequestID string `json:"request_id,omitempty"` - // Deprecated: MA work item 不再返回 ttl_seconds。 - TTLSeconds int `json:"ttl_seconds,omitempty"` - // Deprecated: MA work item 不再返回 lease_seconds。 - LeaseSeconds int `json:"lease_seconds,omitempty"` - // Deprecated: 使用 LatestHeartbeatAt。 - LastHeartbeat string `json:"last_heartbeat,omitempty"` -} - -// SessionIDValue 返回 work item 对应的 session id。 -func (w WorkItem) SessionIDValue() string { - if w.SessionID != "" { - return w.SessionID - } - if w.Data.SessionID != "" { - return w.Data.SessionID - } - if w.Data.ID != "" && (w.Data.Type == "" || w.Data.Type == "session") { - return w.Data.ID - } - return "" -} - -// LeaseTTLSeconds 返回 work item 的 lease TTL,兼容旧字段 lease_seconds。 -func (w WorkItem) LeaseTTLSeconds() int { - if w.TTLSeconds > 0 { - return w.TTLSeconds - } - return w.LeaseSeconds -} - -// LatestHeartbeatValue 返回 work item 的 heartbeat CAS 值,兼容旧字段 last_heartbeat。 -func (w WorkItem) LatestHeartbeatValue() string { - if w.LatestHeartbeatAt != "" { - return w.LatestHeartbeatAt - } - return w.LastHeartbeat -} +type WorkItem = environment.WorkItem // WorkTag 是 work item 关联的火山资源标签。 -type WorkTag struct { - Key string `json:"key"` - Value string `json:"value,omitempty"` -} +type WorkTag = environment.VolcTag // WorkData 是 work item 的业务载荷。 -type WorkData struct { - Type string `json:"type,omitempty"` - ID string `json:"id,omitempty"` - // Deprecated: 使用 ID。 - SessionID string `json:"session_id,omitempty"` -} +type WorkData = environment.WorkData // Session 是 worker 启动一个 session 所需的最小配置快照。 type Session struct {