Skip to content

Commit d0557ee

Browse files
chore: sync upstream main into coder_2_33 (#42)
1 parent a2a3f21 commit d0557ee

111 files changed

Lines changed: 6526 additions & 2842 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎.github/CODEOWNERS‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
* @andreynering

‎.github/workflows/build.yml‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,8 @@
11
name: build
2-
on: [push, pull_request]
2+
on:
3+
push:
4+
branches: [main]
5+
pull_request:
36

47
jobs:
58
govulncheck:

‎.github/workflows/labeler.yml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@ on:
1414

1515
permissions:
1616
issues: write
17+
pull-requests: write
1718
contents: read
1819

1920
jobs:
@@ -26,5 +27,4 @@ jobs:
2627
enable-versioned-regex: 0
2728
include-title: 1
2829
include-body: 0
29-
repo-token: ${{ secrets.PERSONAL_ACCESS_TOKEN }}
3030
issue-number: ${{ github.event.inputs.issue-number || github.event.issue.number || github.event.pull_request.number }}

‎.helix/ignore‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
providertests/testdata

‎agent.go‎

Lines changed: 117 additions & 60 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
564586
func 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.
11941240
func 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

Comments
 (0)