@@ -11,6 +11,7 @@ import (
1111 "slices"
1212 "sync"
1313
14+ "charm.land/fantasy/jsonrepair"
1415 "charm.land/fantasy/schema"
1516 "github.com/charmbracelet/x/exp/slice"
1617)
@@ -146,6 +147,7 @@ type agentSettings struct {
146147 providerDefinedTools []ProviderDefinedTool
147148 executableProviderTools []ExecutableProviderTool
148149 tools []AgentTool
150+ toolChoice * ToolChoice
149151 maxRetries * int
150152
151153 model LanguageModel
@@ -162,12 +164,13 @@ type AgentCall struct {
162164 Files []FilePart `json:"files"`
163165 Messages []Message `json:"messages"`
164166 MaxOutputTokens * int64
165- Temperature * float64 `json:"temperature"`
166- TopP * float64 `json:"top_p"`
167- TopK * int64 `json:"top_k"`
168- PresencePenalty * float64 `json:"presence_penalty"`
169- FrequencyPenalty * float64 `json:"frequency_penalty"`
170- ActiveTools []string `json:"active_tools"`
167+ Temperature * float64 `json:"temperature"`
168+ TopP * float64 `json:"top_p"`
169+ TopK * int64 `json:"top_k"`
170+ PresencePenalty * float64 `json:"presence_penalty"`
171+ FrequencyPenalty * float64 `json:"frequency_penalty"`
172+ ActiveTools []string `json:"active_tools"`
173+ ToolChoice * ToolChoice `json:"tool_choice"`
171174 ProviderOptions ProviderOptions
172175 OnRetry OnRetryCallback
173176 MaxRetries * int
@@ -252,12 +255,13 @@ type AgentStreamCall struct {
252255 Files []FilePart `json:"files"`
253256 Messages []Message `json:"messages"`
254257 MaxOutputTokens * int64
255- Temperature * float64 `json:"temperature"`
256- TopP * float64 `json:"top_p"`
257- TopK * int64 `json:"top_k"`
258- PresencePenalty * float64 `json:"presence_penalty"`
259- FrequencyPenalty * float64 `json:"frequency_penalty"`
260- ActiveTools []string `json:"active_tools"`
258+ Temperature * float64 `json:"temperature"`
259+ TopP * float64 `json:"top_p"`
260+ TopK * int64 `json:"top_k"`
261+ PresencePenalty * float64 `json:"presence_penalty"`
262+ FrequencyPenalty * float64 `json:"frequency_penalty"`
263+ ActiveTools []string `json:"active_tools"`
264+ ToolChoice * ToolChoice `json:"tool_choice"`
261265 Headers map [string ]string
262266 ProviderOptions ProviderOptions
263267 OnRetry OnRetryCallback
@@ -335,6 +339,7 @@ func (a *agent) prepareCall(call AgentCall) AgentCall {
335339 call .PresencePenalty = cmp .Or (call .PresencePenalty , a .settings .presencePenalty )
336340 call .FrequencyPenalty = cmp .Or (call .FrequencyPenalty , a .settings .frequencyPenalty )
337341 call .MaxRetries = cmp .Or (call .MaxRetries , a .settings .maxRetries )
342+ call .ToolChoice = cmp .Or (call .ToolChoice , a .settings .toolChoice )
338343
339344 if len (call .StopWhen ) == 0 && len (a .settings .stopWhen ) > 0 {
340345 call .StopWhen = a .settings .stopWhen
@@ -383,6 +388,9 @@ func (a *agent) Generate(ctx context.Context, opts AgentCall) (*AgentResult, err
383388 stepSystemPrompt := a .settings .systemPrompt
384389 stepActiveTools := opts .ActiveTools
385390 stepToolChoice := ToolChoiceAuto
391+ if opts .ToolChoice != nil {
392+ stepToolChoice = * opts .ToolChoice
393+ }
386394 disableAllTools := false
387395 stepTools := a .settings .tools
388396 if opts .PrepareStep != nil {
@@ -485,7 +493,12 @@ func (a *agent) Generate(ctx context.Context, opts AgentCall) (*AgentResult, err
485493
486494 toolResults , err := a .executeTools (ctx , stepTools , stepExecProviderTools , stepToolCalls , nil )
487495
488- // Build step content with validated tool calls and tool results. // Provider-executed tool calls are kept as-is.
496+ // If any tool result requested a stop, deliver all results but don't
497+ // request another completion from the model.
498+ stopTurnRequested := hasStopTurn (toolResults )
499+
500+ // Build step content with validated tool calls and tool results.
501+ // Provider-executed tool calls are kept as-is.
489502 stepContent := []Content {}
490503 toolCallIndex := 0
491504 for _ , content := range result .Content {
@@ -523,7 +536,7 @@ func (a *agent) Generate(ctx context.Context, opts AgentCall) (*AgentResult, err
523536 steps = append (steps , stepResult )
524537 shouldStop := isStopConditionMet (opts .StopWhen , steps )
525538
526- if shouldStop || err != nil || len (stepToolCalls ) == 0 || result .FinishReason != FinishReasonToolCalls {
539+ if shouldStop || err != nil || stopTurnRequested || len (stepToolCalls ) == 0 || result .FinishReason != FinishReasonToolCalls {
527540 break
528541 }
529542 }
@@ -561,6 +574,15 @@ func isStopConditionMet(conditions []StopCondition, steps []StepResult) bool {
561574 return false
562575}
563576
577+ func hasStopTurn (results []ToolResultContent ) bool {
578+ for _ , r := range results {
579+ if r .StopTurn {
580+ return true
581+ }
582+ }
583+ return false
584+ }
585+
564586func toResponseMessages (content []Content ) []Message {
565587 var assistantParts []MessagePart
566588 var toolParts []MessagePart
@@ -729,13 +751,15 @@ func (a *agent) executeSingleTool(ctx context.Context, toolMap map[string]AgentT
729751 Error : err ,
730752 }
731753 result .ClientMetadata = toolResult .Metadata
754+ result .StopTurn = toolResult .StopTurn
732755 if toolResultCallback != nil {
733756 _ = toolResultCallback (result )
734757 }
735758 return result , true
736759 }
737760
738761 result .ClientMetadata = toolResult .Metadata
762+ result .StopTurn = toolResult .StopTurn
739763 if toolResult .IsError {
740764 result .Result = ToolResultOutputContentError {
741765 Error : errors .New (toolResult .Content ),
@@ -771,6 +795,7 @@ func (a *agent) Stream(ctx context.Context, opts AgentStreamCall) (*AgentResult,
771795 PresencePenalty : opts .PresencePenalty ,
772796 FrequencyPenalty : opts .FrequencyPenalty ,
773797 ActiveTools : opts .ActiveTools ,
798+ ToolChoice : opts .ToolChoice ,
774799 ProviderOptions : opts .ProviderOptions ,
775800 MaxRetries : opts .MaxRetries ,
776801 OnRetry : opts .OnRetry ,
@@ -801,6 +826,9 @@ func (a *agent) Stream(ctx context.Context, opts AgentStreamCall) (*AgentResult,
801826 stepSystemPrompt := a .settings .systemPrompt
802827 stepActiveTools := call .ActiveTools
803828 stepToolChoice := ToolChoiceAuto
829+ if call .ToolChoice != nil {
830+ stepToolChoice = * call .ToolChoice
831+ }
804832 disableAllTools := false
805833 stepTools := a .settings .tools
806834 // Apply step preparation if provided
@@ -1015,6 +1043,16 @@ func (a *agent) validateAndRepairToolCall(ctx context.Context, toolCall ToolCall
10151043 return * repairedToolCall
10161044 }
10171045 }
1046+ } else {
1047+ // Default repair: try jsonrepair for malformed JSON when no
1048+ // custom repair function is configured.
1049+ if repaired , repairErr := jsonrepair .RepairJSON (toolCall .Input ); repairErr == nil && repaired != toolCall .Input {
1050+ repairedCall := toolCall
1051+ repairedCall .Input = repaired
1052+ if validateErr := a .validateToolCall (repairedCall , availableTools , execProviderTools ); validateErr == nil {
1053+ return repairedCall
1054+ }
1055+ }
10181056 }
10191057
10201058 invalidToolCall := toolCall
@@ -1190,6 +1228,14 @@ func WithProviderDefinedTools(tools ...ProviderTool) AgentOption {
11901228 }
11911229}
11921230
1231+ // WithToolChoice sets the default tool choice for the agent. It is overridden
1232+ // by the ToolChoice on a specific call, and by PrepareStep at the step level.
1233+ func WithToolChoice (choice ToolChoice ) AgentOption {
1234+ return func (s * agentSettings ) {
1235+ s .toolChoice = & choice
1236+ }
1237+ }
1238+
11931239// WithStopConditions sets the stop conditions for the agent.
11941240func WithStopConditions (conditions ... StopCondition ) AgentOption {
11951241 return func (s * agentSettings ) {
@@ -1247,11 +1293,7 @@ func (a *agent) processStepStream(ctx context.Context, stream StreamResponse, op
12471293 toolCall ToolCallContent
12481294 parallel bool
12491295 }
1250- toolChan := make (chan toolExecutionRequest , 10 )
1251- var toolExecutionWg sync.WaitGroup
1252- var toolStateMu sync.Mutex
1253- toolResults := make ([]ToolResultContent , 0 )
1254- var toolExecutionErr error
1296+ var pendingDispatches []toolExecutionRequest
12551297
12561298 // Create a map for quick tool lookup
12571299 toolMap := make (map [string ]AgentTool )
@@ -1264,43 +1306,6 @@ func (a *agent) processStepStream(ctx context.Context, stream StreamResponse, op
12641306 execProviderToolMap [ept .GetName ()] = ept
12651307 }
12661308
1267- // Semaphores for controlling parallelism
1268- parallelSem := make (chan struct {}, 5 )
1269- var sequentialMu sync.Mutex
1270-
1271- // Single coordinator goroutine that dispatches tools
1272- toolExecutionWg .Go (func () {
1273- for req := range toolChan {
1274- if req .parallel {
1275- parallelSem <- struct {}{}
1276- toolExecutionWg .Go (func () {
1277- defer func () { <- parallelSem }()
1278- result , isCriticalError := a .executeSingleTool (ctx , toolMap , execProviderToolMap , req .toolCall , opts .OnToolResult )
1279- toolStateMu .Lock ()
1280- toolResults = append (toolResults , result )
1281- if isCriticalError && toolExecutionErr == nil {
1282- if errorResult , ok := result .Result .(ToolResultOutputContentError ); ok && errorResult .Error != nil {
1283- toolExecutionErr = errorResult .Error
1284- }
1285- }
1286- toolStateMu .Unlock ()
1287- })
1288- } else {
1289- sequentialMu .Lock ()
1290- result , isCriticalError := a .executeSingleTool (ctx , toolMap , execProviderToolMap , req .toolCall , opts .OnToolResult )
1291- toolStateMu .Lock ()
1292- toolResults = append (toolResults , result )
1293- if isCriticalError && toolExecutionErr == nil {
1294- if errorResult , ok := result .Result .(ToolResultOutputContentError ); ok && errorResult .Error != nil {
1295- toolExecutionErr = errorResult .Error
1296- }
1297- }
1298- toolStateMu .Unlock ()
1299- sequentialMu .Unlock ()
1300- }
1301- }
1302- })
1303-
13041309 // Process stream parts
13051310 for part := range stream {
13061311 // Forward all parts to chunk callback
@@ -1475,8 +1480,9 @@ func (a *agent) processStepStream(ctx context.Context, stream StreamResponse, op
14751480 isParallel = tool .Info ().Parallel
14761481 }
14771482
1478- // Send tool call to execution channel
1479- toolChan <- toolExecutionRequest {toolCall : validatedToolCall , parallel : isParallel }
1483+ // Buffer dispatch until stream is fully consumed so that all
1484+ // OnToolCall callbacks complete before any tool result is written.
1485+ pendingDispatches = append (pendingDispatches , toolExecutionRequest {toolCall : validatedToolCall , parallel : isParallel })
14801486
14811487 // Clean up active tool call
14821488 delete (activeToolCalls , part .ID )
@@ -1534,7 +1540,58 @@ func (a *agent) processStepStream(ctx context.Context, stream StreamResponse, op
15341540 }
15351541 }
15361542
1537- // Close the tool execution channel and wait for all executions to complete
1543+ // All tool calls are now collected. Create the execution channel sized to
1544+ // avoid blocking during dispatch, start the coordinator, then flush the batch.
1545+ toolChan := make (chan toolExecutionRequest , len (pendingDispatches ))
1546+ var toolExecutionWg sync.WaitGroup
1547+ var toolStateMu sync.Mutex
1548+ toolResults := make ([]ToolResultContent , 0 , len (pendingDispatches ))
1549+ var toolExecutionErr error
1550+
1551+ // Semaphores for controlling parallelism.
1552+ parallelSem := make (chan struct {}, 5 )
1553+ var sequentialMu sync.Mutex
1554+
1555+ // Single coordinator goroutine that dispatches tools.
1556+ toolExecutionWg .Go (func () {
1557+ for req := range toolChan {
1558+ if req .parallel {
1559+ parallelSem <- struct {}{}
1560+ toolExecutionWg .Go (func () {
1561+ defer func () { <- parallelSem }()
1562+ result , isCriticalError := a .executeSingleTool (ctx , toolMap , execProviderToolMap , req .toolCall , opts .OnToolResult )
1563+ toolStateMu .Lock ()
1564+ toolResults = append (toolResults , result )
1565+ if isCriticalError && toolExecutionErr == nil {
1566+ if errorResult , ok := result .Result .(ToolResultOutputContentError ); ok && errorResult .Error != nil {
1567+ toolExecutionErr = errorResult .Error
1568+ }
1569+ }
1570+ toolStateMu .Unlock ()
1571+ })
1572+ } else {
1573+ sequentialMu .Lock ()
1574+ result , isCriticalError := a .executeSingleTool (ctx , toolMap , execProviderToolMap , req .toolCall , opts .OnToolResult )
1575+ toolStateMu .Lock ()
1576+ toolResults = append (toolResults , result )
1577+ if isCriticalError && toolExecutionErr == nil {
1578+ if errorResult , ok := result .Result .(ToolResultOutputContentError ); ok && errorResult .Error != nil {
1579+ toolExecutionErr = errorResult .Error
1580+ }
1581+ }
1582+ toolStateMu .Unlock ()
1583+ sequentialMu .Unlock ()
1584+ }
1585+ }
1586+ })
1587+
1588+ // Dispatch all buffered tool calls now that every OnToolCall callback has
1589+ // been called, then close and wait.
1590+ for _ , req := range pendingDispatches {
1591+ toolChan <- req
1592+ }
1593+
1594+ // Close the tool execution channel and wait for all executions to complete.
15381595 close (toolChan )
15391596 toolExecutionWg .Wait ()
15401597
@@ -1562,7 +1619,7 @@ func (a *agent) processStepStream(ctx context.Context, stream StreamResponse, op
15621619 }
15631620
15641621 // Determine if we should continue (has tool calls and not stopped)
1565- shouldContinue := len (stepToolCalls ) > 0 && stepFinishReason == FinishReasonToolCalls
1622+ shouldContinue := len (stepToolCalls ) > 0 && stepFinishReason == FinishReasonToolCalls && ! hasStopTurn ( toolResults )
15661623
15671624 return stepExecutionResult {
15681625 StepResult : stepResult ,
0 commit comments