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
116 changes: 72 additions & 44 deletions chat/chat.go
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,9 @@ func (c *Chat) ensureDefaults() {
t := tools.NewTools()
c.Tools = t
}
if c.ToolPolicy == 0 {
c.ToolPolicy = ToolPolicyManual
}
}

func (c *Chat) Session(ctx context.Context, client Client) <-chan StreamEvent {
Expand All @@ -97,17 +100,10 @@ func (c *Chat) Session(ctx context.Context, client Client) <-chan StreamEvent {
// skips nil events
send := func(event StreamEvent) bool {
if event == nil {
if ctx.Err() != nil {
return false
}
return true
}
select {
case result <- event:
return true
case <-ctx.Done():
return false
}
result <- event
return true
}

// event handling
Expand Down Expand Up @@ -144,10 +140,11 @@ func (c *Chat) Session(ctx context.Context, client Client) <-chan StreamEvent {

select {
case <-ctx.Done():
c.handleCompletionEnd(ctx, state, true)
return
case ev, ok := <-state.events:
if !ok {
if !c.handleCompletionEnd(ctx, state) {
if !c.handleCompletionEnd(ctx, state, false) {
return
}
restart = true
Expand Down Expand Up @@ -228,62 +225,93 @@ func (c *Chat) Session(ctx context.Context, client Client) <-chan StreamEvent {
return result
}

func (c *Chat) handleCompletionEnd(ctx context.Context, state *sessionState) (proceed bool) {
func (c *Chat) handleCompletionEnd(ctx context.Context, state *sessionState, stopped bool) (proceed bool) {
proceed = false
// adding collected events to the chat (reasoning, assistant's tokens and tool calls)
if state.thinkingBuilder.Len() != 0 {
c.AppendEvent(NewEventReasoningMessage(state.thinkingBuilder.String()))
ev := NewEventReasoningMessage(state.thinkingBuilder.String())
c.AppendEvent(ev)
state.send(ev)
}
var assistantMsg EventAssistantMessage
if state.builder.Len() != 0 {
c.AppendEvent(NewEventAssistantMessage(state.builder.String()))
assistantMsg = NewEventAssistantMessage(state.builder.String())
c.AppendEvent(assistantMsg)
state.send(assistantMsg)
} else if stopped && state.thinkingBuilder.Len() != 0 {
assistantMsg = NewEventAssistantMessage("")
c.AppendEvent(assistantMsg)
state.send(assistantMsg)
}
for _, call := range state.toolCalls {
for i, call := range state.toolCalls {
if stopped && state.lastToolCall != nil && i == len(state.toolCalls)-1 && &state.toolCalls[i] == state.lastToolCall {
continue
}
c.AppendEvent(call)
}

callAmount := len(state.toolCalls)

// send last tool call if it wasn't sent yet
if !state.flushLastToolCall() {
return
}
state.flushLastToolCall()

// ending current completion
if !state.send(NewEventCompletionEnded(state.toolCalls)) {
return
}
state.send(NewEventCompletionEnded(state.toolCalls))

if callAmount == 0 {
return
if len(state.toolCalls) == 0 {
return false
}
if stopped {
return false
}

// initializing approval waiter
verdicts := state.approval.Wait(ctx, callAmount)
policy := c.ToolPolicy
if policy == ToolPolicyManual {
// initializing approval waiter
verdicts := state.approval.Wait(ctx, len(state.toolCalls))

// processing user verdicts
for verdict := range verdicts {
call := verdict.call
// processing user verdicts
for verdict := range verdicts {
call := verdict.call

var toolMessage EventToolMessage
var toolMessage EventToolMessage

if verdict.Accepted {
callResult, success := c.Tools.Execute(call.Name, call.Content)
toolMessage = NewEventToolMessage(call.CallID, callResult, success)
} else {
msg := c.DeclinedToolMessage
if msg == "" {
msg = DefaultDeclinedToolMessage
if verdict.Accepted {
callResult, success := c.Tools.Execute(call.Name, call.Content)
toolMessage = NewEventToolMessage(call.CallID, callResult, success)
} else {
msg := c.DeclinedToolMessage
if msg == "" {
msg = DefaultDeclinedToolMessage
}
toolMessage = NewEventToolMessage(call.CallID, msg, false)
}
// adding tool message to the chat
c.AppendEvent(toolMessage)

// sending tool message
if !state.send(toolMessage) {
return false
}
toolMessage = NewEventToolMessage(call.CallID, msg, false)
}
// adding tool message to the chat
c.AppendEvent(toolMessage)
return true
}

// sending tool message
if !state.send(toolMessage) {
return
// AutoApprove
if policy == ToolPolicyAutoApprove {
for _, call := range state.toolCalls {
// emit resolved event
if !state.send(NewEventToolCallResolved(call.CallID, true)) {
return false
}
callResult, success := c.Tools.Execute(call.Name, call.Content)
toolMessage := NewEventToolMessage(call.CallID, callResult, success)
c.AppendEvent(toolMessage)
if !state.send(toolMessage) {
return false
}
}
return true
}

return true
// ExitAfter: do not execute tools, just exit the session
return false
}
174 changes: 174 additions & 0 deletions chat/chat_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2258,3 +2258,177 @@ func TestSession_ToolCall_DoubleResolveFromCopies(t *testing.T) {

assert.Equal(t, int32(1), execCount.Load(), "tool must be executed exactly once")
}

// TestToolPolicy_ExitAfter_DoesNotExecuteTools verifies ToolPolicyExitAfter
// does not execute tools automatically.
func TestToolPolicy_ExitAfter_DoesNotExecuteTools(t *testing.T) {
chat := &Chat{
Messages: NewMessages(),
Tools: tools.NewTools(),
ToolPolicy: ToolPolicyExitAfter,
}

tool, err := tools.NewTool("test", "Test tool", func(input map[string]string) (string, error) {
t.Error("tool should not be executed in ExitAfter mode")
return "", nil
})
if err != nil {
t.Fatalf("Failed to create tool: %v", err)
}
chat.Tools.Add(tool)

client := &MockClient{
StreamingEvents: []StreamEvent{
NewEventToolCall("call-1", "test", `{"input":{"a":"b"}}`),
},
}

ctx, cancel := context.WithCancel(context.Background())
events := chat.Session(ctx, client)

// Drain until CompletionEnded
for e := range events {
if _, ok := e.(EventCompletionEnded); ok {
cancel()
break
}
}
}

// TestToolPolicy_AutoApprove_ExecutesTools verifies ToolPolicyAutoApprove
// executes tools automatically.
func TestToolPolicy_AutoApprove_ExecutesTools(t *testing.T) {
chat := &Chat{
Messages: NewMessages(),
Tools: tools.NewTools(),
ToolPolicy: ToolPolicyAutoApprove,
}

execCount := int32(0)
tool, err := tools.NewTool("test", "Test tool", func(input map[string]string) (string, error) {
atomic.AddInt32(&execCount, 1)
return "result", nil
})
if err != nil {
t.Fatalf("Failed to create tool: %v", err)
}
chat.Tools.Add(tool)

client := &MockClient{
StreamingEvents: []StreamEvent{
NewEventToolCall("call-1", "test", `{"input":{"a":"b"}}`),
},
}

ctx, cancel := context.WithCancel(context.Background())
events := chat.Session(ctx, client)

// Drain until CompletionEnded
for e := range events {
if _, ok := e.(EventCompletionEnded); ok {
cancel()
break
}
}

if atomic.LoadInt32(&execCount) < 1 {
t.Fatalf("expected tool to be executed at least once, got %d", atomic.LoadInt32(&execCount))
}
}

// TestEventToolCall_Execute verifies standalone tool execution works.
func TestEventToolCall_Execute(t *testing.T) {
chat := &Chat{
Messages: NewMessages(),
Tools: tools.NewTools(),
}

tool, err := tools.NewTool("test", "Test tool", func(input map[string]string) (string, error) {
return "executed: " + input["x"], nil
})
if err != nil {
t.Fatalf("Failed to create tool: %v", err)
}
chat.Tools.Add(tool)

call := NewEventToolCall("call-1", "test", `{"input":{"x":"1"}}`)
msg := call.Execute(chat.Tools)

if msg.CallID != "call-1" {
t.Errorf("expected call_id call-1, got %s", msg.CallID)
}
if msg.Content != "executed: 1" {
t.Errorf("expected executed content, got '%s'", msg.Content)
}
if !msg.Success {
t.Error("expected success=true")
}
}

// TestSession_CancelSendsCompletionEnded verifies that EventCompletionEnded
// is always delivered to the consumer even when ctx is cancelled mid-generation.
func TestSession_CancelSendsCompletionEnded(t *testing.T) {
chat := &Chat{
Messages: NewMessages(),
Tools: &tools.Tools{},
}

client := &MockClient{
StreamingEvents: []StreamEvent{
NewEventToken("Hello "),
NewEventToken("World"),
},
}

ctx, cancel := context.WithCancel(context.Background())
events := chat.Session(ctx, client)

var received []StreamEvent
// Collect first few events, then cancel
for i := 0; i < 2; i++ {
e, ok := <-events
if !ok {
break
}
received = append(received, e)
}
cancel()

// Drain remaining events
for e := range events {
received = append(received, e)
}

var completionEndedFound bool
for _, e := range received {
if _, ok := e.(EventCompletionEnded); ok {
completionEndedFound = true
break
}
}
if !completionEndedFound {
t.Error("Expected EventCompletionEnded to be sent on cancellation, but it was not received")
}

var completionStartFound bool
for _, e := range received {
if _, ok := e.(EventCompletionStart); ok {
completionStartFound = true
break
}
}
if !completionStartFound {
t.Error("Expected EventCompletionStart to be sent")
}

messages := chat.Messages.Snapshot()
var assistant string
for _, msg := range messages {
if m, ok := msg.(EventAssistantMessage); ok {
assistant = m.Content
}
}
if assistant == "" {
t.Error("Expected partial assistant message to be saved to history")
}
}
14 changes: 14 additions & 0 deletions chat/events.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@ import (
"encoding/json"
"errors"
"sync/atomic"

"github.com/x2d7/interlude/chat/tools"
)

// eventType represents the type of event
Expand Down Expand Up @@ -90,6 +92,18 @@ func (e *EventToolCall) Resolve(accept bool) error {

func (e EventToolCall) getType() eventType { return eventToolCall }

// Execute runs the tool call against a Tools registry and returns the result.
// It can be called independently, outside of the session flow. For ToolPolicyManual
// and ToolPolicyAutoApprove modes, this method is NOT needed — tools are executed
// automatically by the session. Use this only when you want to execute a tool call
// manually (e.g., in ToolPolicyExitAfter mode where the session has ended, or for
// ad-hoc execution). The returned EventToolMessage is NOT added to chat history
// automatically — you must call chat.AppendEvent() yourself if you want to persist it.
func (e EventToolCall) Execute(t *tools.Tools) EventToolMessage {
result, success := t.Execute(e.Name, e.Content)
return NewEventToolMessage(e.CallID, result, success)
}

// NewEventToolCall creates a new EventToolCall
func NewEventToolCall(callID, name string, arguments string) EventToolCall {
return EventToolCall{
Expand Down
Loading
Loading