Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 9 additions & 8 deletions arkruntime/lib/environments/heartbeat.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand All @@ -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
}
Expand All @@ -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) {
Expand Down Expand Up @@ -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)
Expand Down
63 changes: 19 additions & 44 deletions arkruntime/lib/environments/integration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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 {
Expand All @@ -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,
Expand All @@ -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)
Expand Down Expand Up @@ -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)
}
}

Expand Down Expand Up @@ -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 {
Expand All @@ -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
},
Expand All @@ -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 {
Expand Down Expand Up @@ -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)
Expand Down
7 changes: 4 additions & 3 deletions arkruntime/lib/environments/poller.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand Down Expand Up @@ -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 的竞争操作。失败时无法证明当前
Expand Down Expand Up @@ -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
Expand All @@ -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)
}
Expand Down
56 changes: 24 additions & 32 deletions arkruntime/lib/environments/poller_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand All @@ -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
Expand Down Expand Up @@ -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",
Expand All @@ -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",
Expand All @@ -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())
}
Expand All @@ -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",
Expand All @@ -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",
Expand All @@ -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",
Expand Down
Loading
Loading