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
216 changes: 176 additions & 40 deletions arkruntime/selfhosted/session_tool_runner.go
Original file line number Diff line number Diff line change
Expand Up @@ -171,6 +171,7 @@ func (r *SessionToolRunner) run() {
defer close(r.done)
defer close(r.events)
defer r.drainInFlight()
defer r.cancel()
r.err = normalizeRunnerErr(r.runLoop())
}

Expand All @@ -195,6 +196,9 @@ func (r *SessionToolRunner) runLoop() error {
pendingAsk: map[string]Event{},
confirmations: map[string]Event{},
externalTools: map[string]Event{},
toolUseEvents: map[string]Event{},
scheduled: map[string]bool{},
executionDone: make(chan toolExecutionResult, sessionRunnerResultsBuffer),
sessionToolUses: map[string]bool{},
toolUsesSinceStatus: map[string]bool{},
blockingEventIDs: map[string]bool{},
Expand Down Expand Up @@ -231,16 +235,22 @@ func (r *SessionToolRunner) runLoop() error {
}

type toolRunnerState struct {
runner *SessionToolRunner
page string
processed map[string]bool
seen map[string]bool
answered map[string]bool
pendingResults map[string]Event
recoveredResults map[string]bool
pendingAsk map[string]Event
confirmations map[string]Event
externalTools map[string]Event
runner *SessionToolRunner
page string
processed map[string]bool
seen map[string]bool
answered map[string]bool
pendingResults map[string]Event
recoveredResults map[string]bool
pendingAsk map[string]Event
confirmations map[string]Event
externalTools map[string]Event
toolUseEvents map[string]Event
// Tool 按 CMA runner 的语义串行执行,事件消费通过完成通道保持非阻塞。
scheduled map[string]bool
executionQueue []pendingToolEvent
activeExecution *activeToolExecution
executionDone chan toolExecutionResult
sessionToolUses map[string]bool
toolUsesSinceStatus map[string]bool
blockingEventIDs map[string]bool
Expand All @@ -250,8 +260,19 @@ type toolRunnerState struct {
}

type pendingToolEvent struct {
event Event
custom bool
event Event
custom bool
confirmation string
}

type activeToolExecution struct {
pendingToolEvent
cancel context.CancelFunc
}

type toolExecutionResult struct {
pendingToolEvent
result toolset.Result
}

func (s *toolRunnerState) consumeStreamLoop(ctx context.Context, streamer EventStreamer) error {
Expand Down Expand Up @@ -352,6 +373,10 @@ func (s *toolRunnerState) consumeStream(ctx context.Context, stream *EventStream
select {
case <-ctx.Done():
return ctx.Err()
case result := <-s.executionDone:
if err := s.finishToolExecution(ctx, result); err != nil {
return err
}
case <-timerC:
if s.idleExpired() {
return ErrIdleTimeout
Expand Down Expand Up @@ -441,6 +466,11 @@ func (s *toolRunnerState) sleepOrIdle(ctx context.Context, d time.Duration) erro
case <-ctx.Done():
timer.Stop()
return ctx.Err()
case result := <-s.executionDone:
timer.Stop()
if err := s.finishToolExecution(ctx, result); err != nil {
return err
}
case <-timer.C:
}
}
Expand Down Expand Up @@ -539,6 +569,10 @@ func (s *toolRunnerState) consumeList(ctx context.Context) error {
select {
case <-ctx.Done():
return ctx.Err()
case result := <-s.executionDone:
if err := s.finishToolExecution(ctx, result); err != nil {
return err
}
case <-timerC:
if s.idleExpired() {
return ErrIdleTimeout
Expand All @@ -551,6 +585,7 @@ func (s *toolRunnerState) consumeList(ctx context.Context) error {
func (s *toolRunnerState) processListedEvents(ctx context.Context, events []Event, reconcile bool) error {
var pending []pendingToolEvent
pendingIDs := map[string]bool{}
replayedToolUses := map[string]bool{}
touchedIdle := false
lastWasEndTurn := false
for _, event := range events {
Expand All @@ -559,24 +594,32 @@ func (s *toolRunnerState) processListedEvents(ctx context.Context, events []Even
continue
}
s.observeSessionState(event)
if event.Type != EventTypeUserToolConfirmation {
if seenNow && event.Type != EventTypeUserToolConfirmation {
touchedIdle = true
lastWasEndTurn = event.Type == EventTypeSessionStatusIdle &&
event.StopReasonType() == SessionStopReasonEndTurn
}
switch event.Type {
case EventTypeUserToolConfirmation:
s.recordConfirmation(event)
case EventTypeUserInterrupt:
if reconcile {
s.handleInterruptForCalls(event, replayedToolUses)
} else {
s.handleInterrupt(event)
}
case EventTypeUserToolResult, EventTypeUserCustomToolResult:
s.markAnswered(toolResultCallID(event))
case EventTypeAgentToolUse:
callID := toolUseCallID(event)
replayedToolUses[callID] = true
if !pendingIDs[callID] {
pending = append(pending, pendingToolEvent{event: event})
pendingIDs[callID] = true
}
case EventTypeAgentCustomToolUse:
callID := toolUseCallID(event)
replayedToolUses[callID] = true
if !pendingIDs[callID] {
pending = append(pending, pendingToolEvent{event: event, custom: true})
pendingIDs[callID] = true
Expand All @@ -602,11 +645,7 @@ func (s *toolRunnerState) processListedEvents(ctx context.Context, events []Even
return err
}
if touchedIdle && lastWasEndTurn {
if s.hasUnblockedOutstandingTool(pending) {
s.disarmIdle()
} else {
s.armIdle()
}
s.armIdle()
}
return nil
}
Expand Down Expand Up @@ -666,7 +705,7 @@ func (s *toolRunnerState) maybeArmPendingIdle() {
}

func (s *toolRunnerState) hasIdleBlockers() bool {
return len(s.pendingAsk) > 0 || len(s.pendingResults) > 0 || len(s.externalTools) > 0
return len(s.pendingAsk) > 0 || len(s.pendingResults) > 0 || len(s.externalTools) > 0 || len(s.scheduled) > 0
}

func (s *toolRunnerState) idleExpired() bool {
Expand All @@ -678,6 +717,8 @@ func (s *toolRunnerState) handleEvent(ctx context.Context, event Event) error {
case EventTypeUserToolConfirmation:
s.recordConfirmation(event)
return s.releaseConfirmedToolUses(ctx)
case EventTypeUserInterrupt:
s.handleInterrupt(event)
case EventTypeUserToolResult, EventTypeUserCustomToolResult:
s.markAnswered(toolResultCallID(event))
case EventTypeAgentToolUse:
Expand All @@ -695,6 +736,9 @@ func (s *toolRunnerState) markEventSeen(event Event) bool {
if key == "" {
key = toolUseCallID(event)
}
if key == "" && event.Type == EventTypeUserInterrupt {
key = fmt.Sprintf("interrupt:%s:%s", event.ProcessedAt, event.SessionThreadID)
}
if key == "" {
return true
}
Expand All @@ -715,6 +759,7 @@ func (s *toolRunnerState) markAnswered(callID string) {
delete(s.recoveredResults, callID)
delete(s.pendingAsk, callID)
delete(s.externalTools, callID)
delete(s.toolUseEvents, callID)
s.maybeArmPendingIdle()
}

Expand Down Expand Up @@ -745,26 +790,10 @@ func (s *toolRunnerState) releaseConfirmedToolUses(ctx context.Context) error {
return nil
}

func (s *toolRunnerState) hasUnblockedOutstandingTool(pending []pendingToolEvent) bool {
for _, toolEvent := range pending {
callID := toolUseCallID(toolEvent.event)
if callID == "" || s.isAnswered(callID) || !s.shouldHandleToolUse(callID) {
continue
}
if _, ok := s.pendingAsk[callID]; ok {
continue
}
if _, ok := s.pendingResults[callID]; ok {
continue
}
return true
}
return false
}

func (s *toolRunnerState) handleToolUse(ctx context.Context, event Event, custom bool) error {
s.ensureRecoveryMaps()
callID := toolUseCallID(event)
if callID == "" || s.isAnswered(callID) {
if callID == "" || s.isAnswered(callID) || s.scheduled[callID] {
return nil
}
if pending := s.pendingResults[callID]; pending.ID != "" {
Expand Down Expand Up @@ -795,9 +824,40 @@ func (s *toolRunnerState) handleToolUse(ctx context.Context, event Event, custom
return s.sendResult(ctx, callID, event, custom, "", decision.Result)
}
}
var result toolset.Result
s.scheduled[callID] = true
s.executionQueue = append(s.executionQueue, pendingToolEvent{
event: event,
custom: custom,
confirmation: confirmation,
})
s.startNextToolExecution(ctx)
return nil
}

func (s *toolRunnerState) startNextToolExecution(ctx context.Context) {
if s.activeExecution != nil {
return
}
for len(s.executionQueue) > 0 {
pending := s.executionQueue[0]
s.executionQueue = s.executionQueue[1:]
callID := toolUseCallID(pending.event)
if s.isAnswered(callID) {
delete(s.scheduled, callID)
continue
}
toolCtx, cancel := context.WithCancel(ctx)
s.activeExecution = &activeToolExecution{pendingToolEvent: pending, cancel: cancel}
go s.executeTool(toolCtx, pending)
return
}
}

func (s *toolRunnerState) executeTool(ctx context.Context, pending pendingToolEvent) {
event := pending.event
input := json.RawMessage(event.Input)
if custom {
var result toolset.Result
if pending.custom {
tool := s.runner.opts.CustomTools[event.Name]
result = s.executeWithTimeout(ctx, event, func(toolCtx context.Context) toolset.Result {
return tool.Execute(toolCtx, input)
Expand All @@ -807,7 +867,71 @@ func (s *toolRunnerState) handleToolUse(ctx context.Context, event Event, custom
return s.runner.opts.Tools.Execute(toolCtx, event.Name, input)
})
}
return s.postResult(ctx, event, custom, callID, result, confirmation)
select {
case s.executionDone <- toolExecutionResult{pendingToolEvent: pending, result: result}:
case <-s.runner.ctx.Done():
}
}

func (s *toolRunnerState) finishToolExecution(ctx context.Context, result toolExecutionResult) error {
callID := toolUseCallID(result.event)
if active := s.activeExecution; active != nil && toolUseCallID(active.event) == callID {
active.cancel()
s.activeExecution = nil
}
delete(s.scheduled, callID)
if !s.isAnswered(callID) {
if err := s.postResult(ctx, result.event, result.custom, callID, result.result, result.confirmation); err != nil {
return err
}
}
s.startNextToolExecution(ctx)
return nil
}

func (s *toolRunnerState) handleInterrupt(event Event) {
s.handleInterruptForCalls(event, nil)
}

func (s *toolRunnerState) handleInterruptForCalls(event Event, eligible map[string]bool) {
threadID := event.SessionThreadID
matches := func(callID string, toolEvent Event) bool {
if eligible != nil && !eligible[callID] {
return false
}
return threadID == "" || toolEvent.SessionThreadID == threadID
}
if active := s.activeExecution; active != nil && matches(toolUseCallID(active.event), active.event) {
active.cancel()
s.settleInterruptedToolUse(toolUseCallID(active.event))
}
queued := s.executionQueue[:0]
for _, pending := range s.executionQueue {
if matches(toolUseCallID(pending.event), pending.event) {
s.settleInterruptedToolUse(toolUseCallID(pending.event))
continue
}
queued = append(queued, pending)
}
s.executionQueue = queued
for callID, toolEvent := range s.toolUseEvents {
if !s.isAnswered(callID) && matches(callID, toolEvent) {
s.settleInterruptedToolUse(callID)
}
}
}

func (s *toolRunnerState) settleInterruptedToolUse(callID string) {
if callID == "" || s.isAnswered(callID) {
return
}
delete(s.scheduled, callID)
s.markAnswered(callID)
if discarder, ok := s.runner.opts.ResultStore.(ToolResultStoreDiscarder); ok {
if err := discarder.Discard(callID); err != nil {
s.runner.logger.Warn("discard interrupted tool result failed", "tool_use_id", callID, "err", err)
}
}
}

func (s *toolRunnerState) ownsTool(event Event, custom bool) bool {
Expand Down Expand Up @@ -1012,6 +1136,9 @@ func (s *toolRunnerState) observeSessionState(event Event) {
callID := toolUseCallID(event)
if callID != "" {
s.sessionToolUses[callID] = true
if !s.isAnswered(callID) {
s.toolUseEvents[callID] = event
}
s.toolUsesSinceStatus[callID] = true
}
case EventTypeSessionStatusIdle:
Expand Down Expand Up @@ -1043,6 +1170,15 @@ func (s *toolRunnerState) ensureRecoveryMaps() {
if s.blockingEventIDs == nil {
s.blockingEventIDs = map[string]bool{}
}
if s.toolUseEvents == nil {
s.toolUseEvents = map[string]Event{}
}
if s.scheduled == nil {
s.scheduled = map[string]bool{}
}
if s.executionDone == nil {
s.executionDone = make(chan toolExecutionResult, sessionRunnerResultsBuffer)
}
}

func (s *toolRunnerState) reconcileRecoveredResults() {
Expand Down
Loading
Loading