diff --git a/internal/temporal/agent_workflow.go b/internal/temporal/agent_workflow.go index b552960..86f1273 100644 --- a/internal/temporal/agent_workflow.go +++ b/internal/temporal/agent_workflow.go @@ -18,6 +18,12 @@ const ( // Workflow history forever. maxWorkflowToolRounds = 20 + // Provider continuations are successful Messages API responses, not + // transport retries. Bound them independently so a provider that never + // reaches a natural stop cannot grow Workflow history forever. + maxPauseTurnContinuations = 5 + maxOutputContinuations = 3 + // maxModelRequestAttempts bounds provider-level retries for one immutable // model request. Infrastructure failures remain Activity errors and retain // Temporal's unbounded recovery policy. @@ -68,6 +74,18 @@ const ( contextCompactionEventChangeID = "thread-context-compaction-event" contextCompactionEventVersion = 1 + + providerStopReasonChangeID = "provider-stop-reason-state-machine" + providerStopReasonVersion = 1 +) + +type providerResponseDisposition uint8 + +const ( + providerResponseComplete providerResponseDisposition = iota + providerResponseExecuteTools + providerResponseContinuePause + providerResponseContinueOutput ) // runWorkflowTurn owns the plan-act-observe loop in deterministic Workflow @@ -164,6 +182,12 @@ func runWorkflowTurnInternal( workflow.DefaultVersion, contextCompactionEventVersion, ) == contextCompactionEventVersion + providerStopReasons := workflow.GetVersion( + actx, + providerStopReasonChangeID, + workflow.DefaultVersion, + providerStopReasonVersion, + ) == providerStopReasonVersion initialOutput := append([]domain.EventDraft(nil), prepared.PreludeEvents...) if contextCompactionEvents && prepared.ContextProjection.Compacted { initialOutput = append(initialOutput, domain.EventDraft{ @@ -221,6 +245,11 @@ func runWorkflowTurnInternal( } outcomeIteration := 0 outcomeFinished := false + pauseTurnContinuations := 0 + outputContinuations := 0 + pauseChainActive := false + var pauseMessagesBase []domain.Message + var pauseTranscriptBase []domain.Message for round := 0; round < maxRounds; round++ { // A later model request must never overtake completed public progress // from the preceding tool/outcome round in PostgreSQL receipt order. @@ -328,7 +357,35 @@ func runWorkflowTurnInternal( return RunTurnResult{}, err } } + var toolUses []domain.ContentBlock + for _, block := range called.Response.Content { + if block.Type == "tool_use" { + toolUses = append(toolUses, block) + } + } + disposition := providerResponseComplete + providerFailure := "" + if providerStopReasons { + disposition, providerFailure = classifyProviderResponse( + called.Response.StopReason, + len(toolUses), + ) + } + if prepared.UsesProviderTranscript { + if providerStopReasons && disposition == providerResponseContinuePause { + if pauseChainActive { + turn.transcriptDelta = append( + []domain.Message(nil), + pauseTranscriptBase..., + ) + } else { + pauseTranscriptBase = append( + []domain.Message(nil), + turn.transcriptDelta..., + ) + } + } turn.transcriptDelta = agentruntime.AppendMerging( turn.transcriptDelta, []domain.Message{{ @@ -337,6 +394,9 @@ func runWorkflowTurnInternal( }}, ) for _, planned := range called.ToolSteps { + if providerStopReasons && disposition != providerResponseExecuteTools { + break + } providerID := planned.ProviderToolUseID if providerID == "" { providerID = planned.ToolUseEventID @@ -361,6 +421,9 @@ func runWorkflowTurnInternal( ) } } + if providerStopReasons && disposition != providerResponseContinuePause { + pauseChainActive = false + } if called.ThinkingEventID != "" { turn.output = append(turn.output, domain.EventDraft{ ID: called.ThinkingEventID, Type: domain.EvAgentThinking, @@ -384,10 +447,77 @@ func runWorkflowTurnInternal( }) } - var toolUses []domain.ContentBlock - for _, block := range called.Response.Content { - if block.Type == "tool_use" { - toolUses = append(toolUses, block) + if providerFailure != "" { + if end := modelRequestEndDraft(called, false); end != nil { + turn.output = append(turn.output, *end) + } + if activityOutcome.Interrupted { + return turn.complete(nil) + } + return turn.terminateTyped( + "model_request_failed_error", + providerFailure, + ) + } + if providerStopReasons { + switch disposition { + case providerResponseContinuePause: + if end := modelRequestEndDraft(called, false); end != nil { + turn.output = append(turn.output, *end) + } + if activityOutcome.Interrupted { + return turn.complete(nil) + } + if pauseTurnContinuations >= maxPauseTurnContinuations { + return turn.terminateTyped( + "model_request_failed_error", + "model exceeded the pause_turn continuation limit", + ) + } + pauseTurnContinuations++ + if pauseChainActive { + messages = append([]domain.Message(nil), pauseMessagesBase...) + } else { + pauseMessagesBase = append([]domain.Message(nil), messages...) + } + messages = agentruntime.AppendMerging(messages, []domain.Message{{ + Role: domain.RoleAssistant, + Content: called.Response.Content, + }}) + pauseChainActive = true + continue + case providerResponseContinueOutput: + if end := modelRequestEndDraft(called, false); end != nil { + turn.output = append(turn.output, *end) + } + if activityOutcome.Interrupted { + return turn.complete(nil) + } + if outputContinuations >= maxOutputContinuations { + return turn.terminateTyped( + "model_request_failed_error", + "model exceeded the max_tokens continuation limit", + ) + } + outputContinuations++ + recoveryMessage := domain.Message{ + Role: domain.RoleUser, + Content: []domain.ContentBlock{{ + Type: "text", + Text: "Output token limit reached. Continue directly from where you stopped without apologizing or recapping.", + }}, + } + messages = agentruntime.AppendMerging(messages, []domain.Message{ + {Role: domain.RoleAssistant, Content: called.Response.Content}, + recoveryMessage, + }) + if prepared.UsesProviderTranscript { + turn.transcriptDelta = agentruntime.AppendMerging( + turn.transcriptDelta, + []domain.Message{recoveryMessage}, + ) + } + continue } } if len(toolUses) == 0 { @@ -608,11 +738,54 @@ func runWorkflowTurnInternal( }) } - // Reaching the safety bound closes the public turn normally rather than - // allowing unbounded Workflow history. + if providerStopReasons { + return turn.terminateTyped( + "model_request_failed_error", + "agent loop exceeded the per-turn round limit", + ) + } + // Preserve the legacy replay outcome for Workflow histories that crossed + // the safety bound before stop-reason handling was introduced. return turn.complete(nil) } +func classifyProviderResponse( + stopReason string, + toolUseCount int, +) (providerResponseDisposition, string) { + switch stopReason { + case "end_turn", "stop_sequence", "refusal", "model_context_window_exceeded": + if toolUseCount > 0 { + return providerResponseComplete, + "model returned " + stopReason + " with a client tool_use block" + } + return providerResponseComplete, "" + case "tool_use": + if toolUseCount == 0 { + return providerResponseComplete, + "model returned tool_use without a client tool_use block" + } + return providerResponseExecuteTools, "" + case "pause_turn": + if toolUseCount > 0 { + return providerResponseComplete, + "model returned pause_turn with a client tool_use block" + } + return providerResponseContinuePause, "" + case "max_tokens": + if toolUseCount > 0 { + return providerResponseComplete, + "model returned max_tokens with a potentially incomplete client tool_use block" + } + return providerResponseContinueOutput, "" + case "": + return providerResponseComplete, "model response has no stop_reason" + default: + return providerResponseComplete, + "model returned unsupported stop_reason " + strconv.Quote(stopReason) + } +} + // modelRequestSpanIDs are deterministic Workflow-owned operation ids. Owning // them before the interruptible model Activity starts lets an interrupt commit // the terminal span.model_request_end that closes any best-effort preview, even diff --git a/internal/temporal/agent_workflow_stop_reason_test.go b/internal/temporal/agent_workflow_stop_reason_test.go new file mode 100644 index 0000000..b35e463 --- /dev/null +++ b/internal/temporal/agent_workflow_stop_reason_test.go @@ -0,0 +1,361 @@ +package temporal + +import ( + "context" + "encoding/json" + "fmt" + "sync" + "testing" + + "github.com/stretchr/testify/require" + "go.temporal.io/sdk/testsuite" + + "github.com/yanpgwang/managed-agent-go/internal/domain" + "github.com/yanpgwang/managed-agent-go/internal/model" +) + +func TestClassifyProviderResponse(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + stopReason string + toolUses int + want providerResponseDisposition + wantFailure string + }{ + {name: "end turn", stopReason: "end_turn", want: providerResponseComplete}, + {name: "stop sequence", stopReason: "stop_sequence", want: providerResponseComplete}, + {name: "refusal", stopReason: "refusal", want: providerResponseComplete}, + {name: "context limit", stopReason: "model_context_window_exceeded", want: providerResponseComplete}, + {name: "tool use", stopReason: "tool_use", toolUses: 1, want: providerResponseExecuteTools}, + {name: "pause turn", stopReason: "pause_turn", want: providerResponseContinuePause}, + {name: "output limit", stopReason: "max_tokens", want: providerResponseContinueOutput}, + {name: "tool reason without tool", stopReason: "tool_use", wantFailure: "without a client tool_use"}, + {name: "pause with client tool", stopReason: "pause_turn", toolUses: 1, wantFailure: "with a client tool_use"}, + {name: "truncated client tool", stopReason: "max_tokens", toolUses: 1, wantFailure: "potentially incomplete"}, + {name: "final reason with tool", stopReason: "end_turn", toolUses: 1, wantFailure: "with a client tool_use"}, + {name: "missing reason", wantFailure: "no stop_reason"}, + {name: "unknown reason", stopReason: "future_reason", wantFailure: `unsupported stop_reason "future_reason"`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + got, failure := classifyProviderResponse(tt.stopReason, tt.toolUses) + require.Equal(t, tt.want, got) + if tt.wantFailure == "" { + require.Empty(t, failure) + } else { + require.Contains(t, failure, tt.wantFailure) + } + }) + } +} + +func TestWorkflowTurn_ContinuesPauseTurnWithExactProviderContent(t *testing.T) { + var suite testsuite.WorkflowTestSuite + env := suite.NewTestWorkflowEnvironment() + env.RegisterWorkflow(workflowTurnHarness) + + initial := []domain.Message{{ + Role: domain.RoleUser, + Content: []domain.ContentBlock{{ + Type: "text", Text: "research the topic", + }}, + }} + serverToolBlock := domain.ContentBlock{ + Type: "server_tool_use", + Raw: json.RawMessage(`{"type":"server_tool_use","id":"srvtoolu_1","name":"web_search","input":{"query":"topic"}}`), + } + pauseContent := []domain.ContentBlock{ + {Type: "text", Text: "I am still searching."}, + serverToolBlock, + } + secondPauseContent := []domain.ContentBlock{ + {Type: "text", Text: "I am checking the remaining sources."}, + { + Type: "server_tool_use", + Raw: json.RawMessage(`{"type":"server_tool_use","id":"srvtoolu_2","name":"web_search","input":{"query":"topic evidence"}}`), + }, + } + + var mu sync.Mutex + var requests []model.Request + var completed CompleteWorkflowTurnInput + registerWorkflowTurnActivities( + env, + func(context.Context, PrepareTurnInput) (PrepareTurnResult, error) { + return PrepareTurnResult{ + ThreadID: "sthr_pause", + UsesProviderTranscript: true, + TranscriptDelta: initial, + Request: model.Request{ + Model: "test-model", + Messages: initial, + }, + }, nil + }, + func(_ context.Context, in CallModelInput) (CallModelResult, error) { + mu.Lock() + requests = append(requests, in.Request) + call := len(requests) + mu.Unlock() + if call == 1 { + return CallModelResult{ + ModelRequestStartID: in.ModelRequestStartID, + ModelRequestEndID: in.ModelRequestEndID, + MessageEventID: "sevt_pause_message", + Response: model.Response{ + StopReason: "pause_turn", + Content: pauseContent, + }, + }, nil + } + if call == 2 { + return CallModelResult{ + ModelRequestStartID: in.ModelRequestStartID, + ModelRequestEndID: in.ModelRequestEndID, + MessageEventID: "sevt_second_pause_message", + Response: model.Response{ + StopReason: "pause_turn", + Content: secondPauseContent, + }, + }, nil + } + return CallModelResult{ + ModelRequestStartID: in.ModelRequestStartID, + ModelRequestEndID: in.ModelRequestEndID, + MessageEventID: "sevt_pause_done", + Response: model.Response{ + StopReason: "end_turn", + Content: []domain.ContentBlock{{Type: "text", Text: "Research complete."}}, + }, + }, nil + }, + func(context.Context, ExecuteToolInput) (ExecuteToolResult, error) { + t.Fatal("pause_turn must not execute a client tool") + return ExecuteToolResult{}, nil + }, + func(_ context.Context, in CompleteWorkflowTurnInput) (RunTurnResult, error) { + completed = in + return RunTurnResult{Disposition: TurnCompleted}, nil + }, + ) + + env.ExecuteWorkflow(workflowTurnHarness, PrepareTurnInput{ + SessionID: "sesn_pause", TriggerEventID: "sevt_user", + }) + require.NoError(t, env.GetWorkflowError()) + + mu.Lock() + defer mu.Unlock() + require.Len(t, requests, 3) + require.Equal(t, initial, requests[0].Messages) + require.Equal(t, []domain.Message{ + initial[0], + {Role: domain.RoleAssistant, Content: pauseContent}, + }, requests[1].Messages) + require.Equal(t, []domain.Message{ + initial[0], + {Role: domain.RoleAssistant, Content: secondPauseContent}, + }, requests[2].Messages, "a repeated pause replaces the prior continuation content") + require.Equal(t, domain.StatusIdle, completed.Status) + require.Equal(t, []string{ + domain.EvAgentMessage, + domain.EvSpanModelRequestEnd, + domain.EvAgentMessage, + domain.EvSpanModelRequestEnd, + domain.EvAgentMessage, + domain.EvSpanModelRequestEnd, + domain.EvSessionStatusIdle, + }, draftTypes(completed.Output)) + require.Len(t, completed.TranscriptDelta, 2) + require.Equal(t, domain.RoleAssistant, completed.TranscriptDelta[1].Role) + require.Equal(t, append( + append([]domain.ContentBlock(nil), secondPauseContent...), + domain.ContentBlock{Type: "text", Text: "Research complete."}, + ), completed.TranscriptDelta[1].Content) +} + +func TestWorkflowTurn_ContinuesMaxTokensWithInternalRecoveryMessage(t *testing.T) { + var suite testsuite.WorkflowTestSuite + env := suite.NewTestWorkflowEnvironment() + env.RegisterWorkflow(workflowTurnHarness) + + initial := []domain.Message{{ + Role: domain.RoleUser, + Content: []domain.ContentBlock{{Type: "text", Text: "write a long answer"}}, + }} + partial := []domain.ContentBlock{{Type: "text", Text: "First half"}} + + var mu sync.Mutex + var requests []model.Request + var completed CompleteWorkflowTurnInput + registerWorkflowTurnActivities( + env, + func(context.Context, PrepareTurnInput) (PrepareTurnResult, error) { + return PrepareTurnResult{ + UsesProviderTranscript: true, + TranscriptDelta: initial, + Request: model.Request{Model: "test-model", Messages: initial}, + }, nil + }, + func(_ context.Context, in CallModelInput) (CallModelResult, error) { + mu.Lock() + requests = append(requests, in.Request) + call := len(requests) + mu.Unlock() + response := model.Response{ + StopReason: "max_tokens", + Content: partial, + } + messageID := "sevt_partial" + if call == 2 { + response = model.Response{ + StopReason: "end_turn", + Content: []domain.ContentBlock{{Type: "text", Text: "Second half"}}, + } + messageID = "sevt_complete" + } + return CallModelResult{ + ModelRequestStartID: in.ModelRequestStartID, + ModelRequestEndID: in.ModelRequestEndID, + MessageEventID: messageID, + Response: response, + }, nil + }, + func(context.Context, ExecuteToolInput) (ExecuteToolResult, error) { + t.Fatal("max_tokens recovery must not execute a client tool") + return ExecuteToolResult{}, nil + }, + func(_ context.Context, in CompleteWorkflowTurnInput) (RunTurnResult, error) { + completed = in + return RunTurnResult{Disposition: TurnCompleted}, nil + }, + ) + + env.ExecuteWorkflow(workflowTurnHarness, PrepareTurnInput{ + SessionID: "sesn_max_tokens", TriggerEventID: "sevt_user", + }) + require.NoError(t, env.GetWorkflowError()) + + mu.Lock() + defer mu.Unlock() + require.Len(t, requests, 2) + require.Len(t, requests[1].Messages, 3) + require.Equal(t, domain.RoleAssistant, requests[1].Messages[1].Role) + require.Equal(t, partial, requests[1].Messages[1].Content) + require.Equal(t, domain.RoleUser, requests[1].Messages[2].Role) + require.Contains(t, requests[1].Messages[2].Content[0].Text, "Continue directly") + require.Equal(t, domain.StatusIdle, completed.Status) + require.Equal(t, []string{ + domain.EvAgentMessage, + domain.EvSpanModelRequestEnd, + domain.EvAgentMessage, + domain.EvSpanModelRequestEnd, + domain.EvSessionStatusIdle, + }, draftTypes(completed.Output)) + require.Len(t, completed.TranscriptDelta, 4) + require.Equal(t, requests[1].Messages[2], completed.TranscriptDelta[2]) +} + +func TestWorkflowTurn_RejectsContradictoryStopReason(t *testing.T) { + var suite testsuite.WorkflowTestSuite + env := suite.NewTestWorkflowEnvironment() + env.RegisterWorkflow(workflowTurnHarness) + + var completed CompleteWorkflowTurnInput + registerWorkflowTurnActivities( + env, + func(context.Context, PrepareTurnInput) (PrepareTurnResult, error) { + return PrepareTurnResult{Request: model.Request{Model: "test-model"}}, nil + }, + func(_ context.Context, in CallModelInput) (CallModelResult, error) { + return CallModelResult{ + ModelRequestStartID: in.ModelRequestStartID, + ModelRequestEndID: in.ModelRequestEndID, + MessageEventID: "sevt_invalid_stop", + Response: model.Response{ + StopReason: "tool_use", + Content: []domain.ContentBlock{{Type: "text", Text: "invalid response"}}, + }, + }, nil + }, + func(context.Context, ExecuteToolInput) (ExecuteToolResult, error) { + t.Fatal("invalid response must not execute a client tool") + return ExecuteToolResult{}, nil + }, + func(_ context.Context, in CompleteWorkflowTurnInput) (RunTurnResult, error) { + completed = in + return RunTurnResult{Disposition: TurnTerminated}, nil + }, + ) + + env.ExecuteWorkflow(workflowTurnHarness, PrepareTurnInput{ + SessionID: "sesn_invalid_stop", TriggerEventID: "sevt_user", + }) + require.NoError(t, env.GetWorkflowError()) + require.Equal(t, domain.StatusTerminated, completed.Status) + require.Equal(t, []string{ + domain.EvAgentMessage, + domain.EvSpanModelRequestEnd, + domain.EvSessionError, + domain.EvSessionStatusTerminated, + }, draftTypes(completed.Output)) + errorPayload := completed.Output[2].Payload["error"].(map[string]any) + require.Equal(t, "model_request_failed_error", errorPayload["type"]) + require.Contains(t, errorPayload["message"], "without a client tool_use") +} + +func TestWorkflowTurn_BoundsMaxTokensContinuation(t *testing.T) { + var suite testsuite.WorkflowTestSuite + env := suite.NewTestWorkflowEnvironment() + env.RegisterWorkflow(workflowTurnHarness) + + var mu sync.Mutex + calls := 0 + var completed CompleteWorkflowTurnInput + registerWorkflowTurnActivities( + env, + func(context.Context, PrepareTurnInput) (PrepareTurnResult, error) { + return PrepareTurnResult{Request: model.Request{Model: "test-model"}}, nil + }, + func(_ context.Context, in CallModelInput) (CallModelResult, error) { + mu.Lock() + defer mu.Unlock() + calls++ + return CallModelResult{ + ModelRequestStartID: in.ModelRequestStartID, + ModelRequestEndID: in.ModelRequestEndID, + MessageEventID: fmt.Sprintf("sevt_truncated_%d", calls), + Response: model.Response{ + StopReason: "max_tokens", + Content: []domain.ContentBlock{{ + Type: "text", Text: fmt.Sprintf("part %d", calls), + }}, + }, + }, nil + }, + func(context.Context, ExecuteToolInput) (ExecuteToolResult, error) { + t.Fatal("max_tokens recovery must not execute a client tool") + return ExecuteToolResult{}, nil + }, + func(_ context.Context, in CompleteWorkflowTurnInput) (RunTurnResult, error) { + completed = in + return RunTurnResult{Disposition: TurnTerminated}, nil + }, + ) + + env.ExecuteWorkflow(workflowTurnHarness, PrepareTurnInput{ + SessionID: "sesn_max_tokens_bound", TriggerEventID: "sevt_user", + }) + require.NoError(t, env.GetWorkflowError()) + mu.Lock() + require.Equal(t, maxOutputContinuations+1, calls) + mu.Unlock() + require.Equal(t, domain.StatusTerminated, completed.Status) + require.Equal(t, domain.EvSessionError, completed.Output[len(completed.Output)-2].Type) + errorPayload := completed.Output[len(completed.Output)-2].Payload["error"].(map[string]any) + require.Contains(t, errorPayload["message"], "max_tokens continuation limit") +} diff --git a/internal/temporal/agent_workflow_test.go b/internal/temporal/agent_workflow_test.go index 3fad008..115ee7e 100644 --- a/internal/temporal/agent_workflow_test.go +++ b/internal/temporal/agent_workflow_test.go @@ -657,7 +657,7 @@ func TestWorkflowTurn_MixedExecutableAndPendingToolsCommitExecutedTranscriptResu ToolStepID: "tstep_ask", }, }, - Response: model.Response{Content: []domain.ContentBlock{ + Response: model.Response{StopReason: "tool_use", Content: []domain.ContentBlock{ { Type: "tool_use", ToolUseID: "provider_read", ToolName: "read", Input: map[string]any{"path": "a.txt"}, @@ -805,7 +805,7 @@ func TestWorkflowTurn_AmbiguousToolTerminatesHonestly(t *testing.T) { ToolSteps: []PlannedToolStep{{ ToolUseEventID: "sevt_ambiguous", ToolStepID: "tstep_ambiguous", }}, - Response: model.Response{Content: []domain.ContentBlock{{ + Response: model.Response{StopReason: "tool_use", Content: []domain.ContentBlock{{ Type: "tool_use", ToolUseID: "sevt_ambiguous", ToolName: "bash", Input: map[string]any{"command": "side effect"}, }}}, @@ -863,7 +863,7 @@ func TestWorkflowTurn_PermanentToolPreparationTerminatesHonestly(t *testing.T) { ToolUseEventID: "sevt_resource_permanent", ToolStepID: "tstep_resource_permanent", }}, - Response: model.Response{Content: []domain.ContentBlock{{ + Response: model.Response{StopReason: "tool_use", Content: []domain.ContentBlock{{ Type: "tool_use", ToolUseID: "sevt_resource_permanent", ToolName: "bash", Input: map[string]any{"command": "true"}, }}}, @@ -1221,7 +1221,7 @@ func TestWorkflowTurn_MixedBatchExecutesBuiltinAndParksClientAction(t *testing.T {ToolUseEventID: "sevt_builtin", ToolStepID: "tstep_builtin"}, {ToolUseEventID: "sevt_custom", ToolStepID: "tstep_custom"}, }, - Response: model.Response{Content: []domain.ContentBlock{ + Response: model.Response{StopReason: "tool_use", Content: []domain.ContentBlock{ { Type: "tool_use", ToolUseID: "sevt_builtin", ToolName: "bash", Input: map[string]any{"command": "must not run"}, @@ -1451,7 +1451,7 @@ func TestWorkflowTurn_SelfHostedToolResultResumesWithoutServerExecution(t *testi modelInput = in return CallModelResult{ MessageEventID: "sevt_final", - Response: model.Response{Content: []domain.ContentBlock{{ + Response: model.Response{StopReason: "end_turn", Content: []domain.ContentBlock{{ Type: "text", Text: "used client result", }}}, }, nil @@ -1700,7 +1700,7 @@ func TestWorkflowTurn_ResumeExecutionKeepsToolOrdinal(t *testing.T) { ToolUseEventID: "sevt_bash", ToolStepID: "tstep_bash", }}, - Response: model.Response{Content: []domain.ContentBlock{{ + Response: model.Response{StopReason: "tool_use", Content: []domain.ContentBlock{{ Type: "tool_use", ToolUseID: "sevt_bash", ToolName: "bash", diff --git a/internal/temporal/outcome_test.go b/internal/temporal/outcome_test.go index 0249d27..ca81124 100644 --- a/internal/temporal/outcome_test.go +++ b/internal/temporal/outcome_test.go @@ -185,7 +185,8 @@ func TestWorkflowTurnEvaluatesOutcomeAndAccountsUsage(t *testing.T) { ModelRequestStartID: in.ModelRequestStartID, ModelRequestEndID: in.ModelRequestEndID, Response: model.Response{ - Content: []domain.ContentBlock{{Type: "text", Text: "finished report"}}, + StopReason: "end_turn", + Content: []domain.ContentBlock{{Type: "text", Text: "finished report"}}, Usage: domain.TokenUsage{ InputTokens: 10, OutputTokens: 4, Speed: "standard", }, @@ -261,8 +262,9 @@ func TestWorkflowTurnClosesStartedOutcomeEvaluationOnFatalGraderResult(t *testin ModelRequestStartID: in.ModelRequestStartID, ModelRequestEndID: in.ModelRequestEndID, Response: model.Response{ - Content: []domain.ContentBlock{{Type: "text", Text: "report"}}, - Usage: domain.TokenUsage{InputTokens: 5, OutputTokens: 2}, + StopReason: "end_turn", + Content: []domain.ContentBlock{{Type: "text", Text: "report"}}, + Usage: domain.TokenUsage{InputTokens: 5, OutputTokens: 2}, }, }, nil }