Implement interrupt handling for context management

- Added an interrupt controller to manage user interrupts during streaming operations.
- Introduced methods for registering and unregistering contexts in a global registry to facilitate interrupt handling.
- Enhanced the Assistant's Stream method to check for force interrupts and handle them appropriately.
- Updated context creation to include unique IDs for better identification during interrupts.
- Implemented logic for graceful and force interrupts, allowing for flexible response to user signals.
This commit is contained in:
Max 2025-11-19 18:41:31 +08:00
parent c25fbb22d7
commit eff70ab07d
8 changed files with 1583 additions and 13 deletions

View file

@ -16,6 +16,13 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
var err error
// Set up interrupt handler if interrupt controller is available
if ctx.Interrupt != nil {
ctx.Interrupt.SetHandler(func(c *context.Context, signal *context.InterruptSignal) error {
return ast.handleInterrupt(c, signal)
})
}
// Initialize stack and auto-handle completion/failure/restore
_, traceID, done := context.EnterStack(ctx, ast.ID, ctx.Referer)
defer done()
@ -506,3 +513,34 @@ func (ast *Assistant) getUses() *context.Uses {
func (ast *Assistant) WithHistory(ctx *context.Context, messages []context.Message) ([]context.Message, error) {
return messages, nil
}
// handleInterrupt handles the interrupt signal
// This is called by the interrupt listener when a signal is received
func (ast *Assistant) handleInterrupt(ctx *context.Context, signal *context.InterruptSignal) error {
fmt.Printf("=== Interrupt Received ===\n")
fmt.Printf("Assistant: %s\n", ast.ID)
fmt.Printf("Type: %s\n", signal.Type)
fmt.Printf("Messages: %d\n", len(signal.Messages))
fmt.Printf("Timestamp: %d\n", signal.Timestamp)
// Handle based on interrupt type
switch signal.Type {
case context.InterruptForce:
fmt.Println("Force interrupt: stopping current operations immediately...")
// Force interrupt: context is already cancelled in handleSignal
// LLM streaming will detect ctx.Interrupt.Context().Done() and stop
case context.InterruptGraceful:
fmt.Println("Graceful interrupt: will process after current step completes...")
// Graceful interrupt: let current operation complete
// The signal is stored in current/pending, can be checked at checkpoints
}
// TODO: Implement actual interrupt handling logic:
// 1. For graceful: wait for current step, then merge messages and restart
// 2. For force: immediately stop and restart with new messages
// 3. Call Interrupted Hook if configured
// 4. Decide whether to continue, restart, or abort based on Hook response
return nil
}

View file

@ -0,0 +1,376 @@
package assistant_test
import (
stdContext "context"
"fmt"
"testing"
"time"
"github.com/yaoapp/gou/plan"
"github.com/yaoapp/yao/agent/assistant"
"github.com/yaoapp/yao/agent/context"
"github.com/yaoapp/yao/agent/testutils"
"github.com/yaoapp/yao/openapi/oauth/types"
)
// newTestContextWithInterrupt creates a Context with interrupt controller for testing
func newTestContextWithInterrupt(chatID, assistantID string) *context.Context {
ctx := &context.Context{
Context: stdContext.Background(),
ID: fmt.Sprintf("test_ctx_%d", time.Now().UnixNano()),
Space: plan.NewMemorySharedSpace(),
ChatID: chatID,
AssistantID: assistantID,
Connector: "",
Locale: "en-us",
Theme: "light",
Client: context.Client{
Type: "web",
UserAgent: "TestAgent/1.0",
IP: "127.0.0.1",
},
Referer: context.RefererAPI,
Accept: context.AcceptWebCUI,
Route: "/test/route",
Metadata: map[string]interface{}{
"test": "interrupt_test",
},
Authorized: &types.AuthorizedInfo{
Subject: "test-user",
ClientID: "test-client-id",
UserID: "test-user-123",
TeamID: "test-team-456",
TenantID: "test-tenant-789",
SessionID: "test-session-id",
},
}
// Initialize interrupt controller
ctx.Interrupt = context.NewInterruptController()
ctx.Interrupt.SetContextID(ctx.ID)
// Register context globally
if err := context.Register(ctx); err != nil {
panic(fmt.Sprintf("Failed to register context: %v", err))
}
// Start interrupt listener
ctx.Interrupt.Start()
return ctx
}
// TestAgentInterruptGraceful tests graceful interrupt during agent stream
func TestAgentInterruptGraceful(t *testing.T) {
testutils.Prepare(t)
defer testutils.Clean(t)
agent, err := assistant.Get("tests.interrupt")
if err != nil {
t.Skipf("Skipping test: assistant 'tests.interrupt' not found: %v", err)
return
}
t.Run("GracefulInterruptDuringStream", func(t *testing.T) {
// Create context with interrupt support
ctx := newTestContextWithInterrupt("chat-interrupt-graceful", "tests.interrupt")
defer ctx.Release()
// Track handler invocations
handlerInvoked := false
var receivedSignal *context.InterruptSignal
// Override the handler to track invocations
originalHandler := ctx.Interrupt
ctx.Interrupt.SetHandler(func(c *context.Context, signal *context.InterruptSignal) error {
handlerInvoked = true
receivedSignal = signal
t.Logf("✓ Interrupt handler invoked: type=%s, messages=%d", signal.Type, len(signal.Messages))
return nil
})
inputMessages := []context.Message{
{Role: context.RoleUser, Content: "Tell me a long story about artificial intelligence"},
}
// Start streaming in a goroutine
streamDone := make(chan error, 1)
go func() {
_, err := agent.Stream(ctx, inputMessages)
streamDone <- err
}()
// Wait a bit to ensure stream has started
time.Sleep(300 * time.Millisecond)
// Send graceful interrupt signal
signal := &context.InterruptSignal{
Type: context.InterruptGraceful,
Messages: []context.Message{
{Role: context.RoleUser, Content: "Actually, can you make it shorter?"},
},
Timestamp: time.Now().UnixMilli(),
}
err = context.SendInterrupt(ctx.ID, signal)
if err != nil {
t.Logf("Warning: Failed to send interrupt (stream may have completed): %v", err)
} else {
t.Log("✓ Graceful interrupt signal sent")
}
// Wait for stream to complete (with timeout)
select {
case err := <-streamDone:
if err != nil {
t.Logf("Stream completed with error: %v", err)
} else {
t.Log("✓ Stream completed successfully")
}
case <-time.After(10 * time.Second):
t.Log("Stream timeout (expected for real LLM calls)")
}
// Verify handler was invoked if signal was sent
if originalHandler != nil {
time.Sleep(200 * time.Millisecond) // Wait for async handler
if handlerInvoked {
t.Log("✓ Interrupt handler was invoked")
if receivedSignal != nil && receivedSignal.Type == context.InterruptGraceful {
t.Log("✓ Received graceful interrupt signal")
}
}
}
})
}
// TestAgentInterruptForce tests force interrupt during agent stream
func TestAgentInterruptForce(t *testing.T) {
testutils.Prepare(t)
defer testutils.Clean(t)
agent, err := assistant.Get("tests.interrupt")
if err != nil {
t.Skipf("Skipping test: assistant 'tests.interrupt' not found: %v", err)
return
}
t.Run("ForceInterruptDuringStream", func(t *testing.T) {
// Create context with interrupt support
ctx := newTestContextWithInterrupt("chat-interrupt-force", "tests.interrupt")
defer ctx.Release()
// Track handler invocations
handlerInvoked := false
streamInterrupted := false
ctx.Interrupt.SetHandler(func(c *context.Context, signal *context.InterruptSignal) error {
handlerInvoked = true
t.Logf("✓ Interrupt handler invoked: type=%s", signal.Type)
return nil
})
inputMessages := []context.Message{
{Role: context.RoleUser, Content: "Write a very detailed essay about machine learning"},
}
// Start streaming in a goroutine
streamDone := make(chan error, 1)
go func() {
_, err := agent.Stream(ctx, inputMessages)
streamDone <- err
}()
// Wait a bit to ensure stream has started
time.Sleep(300 * time.Millisecond)
// Send force interrupt signal
signal := &context.InterruptSignal{
Type: context.InterruptForce,
Messages: []context.Message{
{Role: context.RoleUser, Content: "Stop! I need something else now."},
},
Timestamp: time.Now().UnixMilli(),
}
err = context.SendInterrupt(ctx.ID, signal)
if err != nil {
t.Logf("Warning: Failed to send interrupt: %v", err)
} else {
t.Log("✓ Force interrupt signal sent")
}
// Wait for stream to complete or be interrupted
select {
case err := <-streamDone:
if err != nil {
// Check if error is due to interrupt
if err.Error() == "force interrupted by user" ||
err.Error() == "interrupted by user" ||
err.Error() == "interrupted by user before stream start" {
streamInterrupted = true
t.Logf("✓ Stream was interrupted: %v", err)
} else {
t.Logf("Stream completed with error: %v", err)
}
} else {
t.Log("Stream completed without error")
}
case <-time.After(10 * time.Second):
t.Log("Stream timeout")
}
// Verify interrupt behavior
time.Sleep(200 * time.Millisecond)
if handlerInvoked {
t.Log("✓ Force interrupt handler was invoked")
}
if streamInterrupted {
t.Log("✓ Stream was interrupted by force signal")
}
})
}
// TestAgentMultipleInterrupts tests multiple interrupts during stream
func TestAgentMultipleInterrupts(t *testing.T) {
testutils.Prepare(t)
defer testutils.Clean(t)
agent, err := assistant.Get("tests.interrupt")
if err != nil {
t.Skipf("Skipping test: assistant 'tests.interrupt' not found: %v", err)
return
}
t.Run("MultipleGracefulInterrupts", func(t *testing.T) {
// Create context with interrupt support
ctx := newTestContextWithInterrupt("chat-interrupt-multiple", "tests.interrupt")
defer ctx.Release()
handlerCallCount := 0
ctx.Interrupt.SetHandler(func(c *context.Context, signal *context.InterruptSignal) error {
handlerCallCount++
t.Logf("✓ Interrupt handler invoked (call %d): %d messages", handlerCallCount, len(signal.Messages))
return nil
})
inputMessages := []context.Message{
{Role: context.RoleUser, Content: "Explain quantum computing in detail"},
}
// Start streaming
streamDone := make(chan error, 1)
go func() {
_, err := agent.Stream(ctx, inputMessages)
streamDone <- err
}()
// Wait for stream to start
time.Sleep(300 * time.Millisecond)
// Send multiple graceful interrupts
for i := 1; i <= 3; i++ {
signal := &context.InterruptSignal{
Type: context.InterruptGraceful,
Messages: []context.Message{
{Role: context.RoleUser, Content: fmt.Sprintf("Additional question %d", i)},
},
Timestamp: time.Now().UnixMilli(),
}
err = context.SendInterrupt(ctx.ID, signal)
if err != nil {
t.Logf("Warning: Failed to send interrupt %d: %v", i, err)
} else {
t.Logf("✓ Sent interrupt %d", i)
}
time.Sleep(100 * time.Millisecond)
}
// Wait for stream to complete
select {
case err := <-streamDone:
if err != nil {
t.Logf("Stream completed with error: %v", err)
}
case <-time.After(10 * time.Second):
t.Log("Stream timeout")
}
// Check if interrupts were received
time.Sleep(300 * time.Millisecond)
pendingCount := ctx.Interrupt.GetPendingCount()
t.Logf("Handler was called %d times, pending count: %d", handlerCallCount, pendingCount)
if handlerCallCount > 0 {
t.Log("✓ Multiple interrupts were processed")
}
})
}
// TestAgentInterruptWithoutStream tests interrupt behavior when no stream is active
func TestAgentInterruptWithoutStream(t *testing.T) {
t.Run("InterruptBeforeStream", func(t *testing.T) {
// Create context with interrupt support
ctx := newTestContextWithInterrupt("chat-interrupt-before", "test-assistant")
defer ctx.Release()
// Send interrupt before starting stream
signal := &context.InterruptSignal{
Type: context.InterruptGraceful,
Messages: []context.Message{
{Role: context.RoleUser, Content: "Early interrupt"},
},
Timestamp: time.Now().UnixMilli(),
}
err := context.SendInterrupt(ctx.ID, signal)
if err != nil {
t.Fatalf("Failed to send interrupt: %v", err)
}
// Wait for signal to be processed
time.Sleep(100 * time.Millisecond)
// Check if signal is in queue
receivedSignal := ctx.Interrupt.Peek()
if receivedSignal == nil {
t.Fatal("Expected interrupt signal to be queued")
}
if receivedSignal.Type != context.InterruptGraceful {
t.Errorf("Expected graceful interrupt, got: %s", receivedSignal.Type)
}
t.Log("✓ Interrupt queued before stream starts")
})
}
// TestAgentInterruptContextCleanup tests cleanup after interrupt
func TestAgentInterruptContextCleanup(t *testing.T) {
t.Run("CleanupAfterInterrupt", func(t *testing.T) {
ctx := newTestContextWithInterrupt("chat-interrupt-cleanup", "test-assistant")
// Send interrupt
signal := &context.InterruptSignal{
Type: context.InterruptGraceful,
Messages: []context.Message{{Role: context.RoleUser, Content: "test"}},
Timestamp: time.Now().UnixMilli(),
}
context.SendInterrupt(ctx.ID, signal)
time.Sleep(100 * time.Millisecond)
// Release context
ctx.Release()
// Try to send interrupt to released context
err := context.SendInterrupt(ctx.ID, signal)
if err == nil {
t.Error("Expected error when sending to released context")
} else {
t.Logf("✓ Correctly rejected interrupt to released context: %v", err)
}
})
}

View file

@ -3,6 +3,7 @@ package context
import (
"context"
"fmt"
"sync"
"time"
jsoniter "github.com/json-iterator/go"
@ -14,6 +15,11 @@ import (
traceTypes "github.com/yaoapp/yao/trace/types"
)
// Global context registry for interrupt management
var (
contextRegistry = &sync.Map{} // map[contextID]*Context
)
// New create a new context
func New(parent context.Context, authorized *types.AuthorizedInfo, chatID, payload string) Context {
@ -24,6 +30,7 @@ func New(parent context.Context, authorized *types.AuthorizedInfo, chatID, paylo
// Validate the client type
ctx := Context{
Context: parent,
ID: generateContextID(), // Generate unique ID for the context
Space: plan.NewMemorySharedSpace(),
ChatID: chatID,
}
@ -68,6 +75,17 @@ func WithTimeout(parent Context, timeout time.Duration) (Context, context.Cancel
// Release the context and clean up all resources including stacks and trace
func (ctx *Context) Release() {
// Unregister from global registry
if ctx.ID != "" {
Unregister(ctx.ID)
}
// Stop interrupt controller
if ctx.Interrupt != nil {
ctx.Interrupt.Stop()
ctx.Interrupt = nil
}
// Complete and release trace if exists
if ctx.trace != nil && ctx.Stack != nil && ctx.Stack.TraceID != "" {
// Mark trace as complete (sends final event)
@ -242,3 +260,60 @@ func (ctx *Context) Map() map[string]interface{} {
return data
}
// Global Registry Functions
// ===================================
// Register registers a context to the global registry
func Register(ctx *Context) error {
if ctx == nil {
return fmt.Errorf("context is nil")
}
if ctx.ID == "" {
return fmt.Errorf("context ID is empty")
}
contextRegistry.Store(ctx.ID, ctx)
return nil
}
// Unregister removes a context from the global registry
func Unregister(contextID string) {
contextRegistry.Delete(contextID)
}
// Get retrieves a context from the global registry by ID
func Get(contextID string) (*Context, error) {
value, ok := contextRegistry.Load(contextID)
if !ok {
return nil, fmt.Errorf("context not found: %s", contextID)
}
ctx, ok := value.(*Context)
if !ok {
return nil, fmt.Errorf("invalid context type")
}
return ctx, nil
}
// SendInterrupt sends an interrupt signal to a context by ID
// This is the main entry point for external interrupt requests
func SendInterrupt(contextID string, signal *InterruptSignal) error {
ctx, err := Get(contextID)
if err != nil {
return err
}
if ctx.Interrupt == nil {
return fmt.Errorf("interrupt controller not initialized for context: %s", contextID)
}
return ctx.Interrupt.SendSignal(signal)
}
// generateContextID generates a unique context ID
func generateContextID() string {
return fmt.Sprintf("ctx_%d", time.Now().UnixNano())
}

285
agent/context/interrupt.go Normal file
View file

@ -0,0 +1,285 @@
package context
import (
"context"
"fmt"
"time"
)
// NewInterruptController creates a new interrupt controller
func NewInterruptController() *InterruptController {
ctrl := &InterruptController{
queue: make(chan *InterruptSignal, 10), // Buffer for 10 interrupts
pending: make([]*InterruptSignal, 0),
}
ctrl.ctx, ctrl.cancel = context.WithCancel(context.Background())
return ctrl
}
// Start starts the interrupt listener goroutine
func (ic *InterruptController) Start() {
if ic.listenerStarted {
return
}
ic.mutex.Lock()
ic.listenerStarted = true
ic.mutex.Unlock()
go ic.listen()
}
// SetHandler sets the handler for interrupt signals
func (ic *InterruptController) SetHandler(handler InterruptHandler) {
if ic == nil {
return
}
ic.handler = handler
}
// SetContextID sets the context ID for retrieving the parent context
func (ic *InterruptController) SetContextID(contextID string) {
if ic == nil {
return
}
ic.contextID = contextID
}
// listen is the main listener goroutine that processes interrupt signals
func (ic *InterruptController) listen() {
for {
select {
case signal := <-ic.queue:
ic.handleSignal(signal)
case <-ic.ctx.Done():
// Context cancelled, stop listening
return
}
}
}
// handleSignal processes an interrupt signal
func (ic *InterruptController) handleSignal(signal *InterruptSignal) {
if signal == nil {
return
}
ic.mutex.Lock()
// If no current interrupt, set it as current
if ic.current == nil {
ic.current = signal
} else {
// If there's already a current interrupt, add to pending queue
ic.pending = append(ic.pending, signal)
}
// For force interrupt, cancel the interrupt context
// This allows LLM streaming and other operations to check and stop
if signal.Type == InterruptForce {
if ic.cancel != nil {
ic.cancel()
// Create a new context for potential future operations
ic.ctx, ic.cancel = context.WithCancel(context.Background())
}
}
ic.mutex.Unlock()
// Call the registered handler if available (outside lock to avoid deadlock)
if ic.handler != nil && ic.contextID != "" {
go func() {
// Retrieve the parent context from global registry
ctx, err := Get(ic.contextID)
if err != nil {
fmt.Printf("Failed to get context for interrupt handler: %v\n", err)
return
}
// Call the handler
if err := ic.handler(ctx, signal); err != nil {
fmt.Printf("Interrupt handler error: %v\n", err)
}
}()
}
}
// Check checks for current interrupt signal (non-blocking)
// Returns the current interrupt and moves to next one if available
func (ic *InterruptController) Check() *InterruptSignal {
if ic == nil {
return nil
}
ic.mutex.Lock()
defer ic.mutex.Unlock()
if ic.current == nil {
return nil
}
// Get current interrupt
signal := ic.current
// Move to next interrupt in queue
if len(ic.pending) > 0 {
ic.current = ic.pending[0]
ic.pending = ic.pending[1:]
} else {
ic.current = nil
}
return signal
}
// CheckWithMerge checks for interrupts and merges all pending messages
// This is useful when multiple interrupts should be handled together
func (ic *InterruptController) CheckWithMerge() *InterruptSignal {
if ic == nil {
return nil
}
ic.mutex.Lock()
defer ic.mutex.Unlock()
if ic.current == nil {
return nil
}
// If there are pending interrupts, merge all messages
if len(ic.pending) > 0 {
// Collect all messages
allMessages := append([]Message{}, ic.current.Messages...)
for _, pending := range ic.pending {
allMessages = append(allMessages, pending.Messages...)
}
// Create merged signal
mergedSignal := &InterruptSignal{
Type: ic.current.Type, // Use first signal's type
Messages: allMessages,
Timestamp: time.Now().UnixMilli(),
Metadata: map[string]interface{}{
"merged": true,
"merged_count": len(ic.pending) + 1,
"original_time": ic.current.Timestamp,
},
}
// Clear all interrupts
ic.current = nil
ic.pending = make([]*InterruptSignal, 0)
return mergedSignal
}
// No pending interrupts, return current
signal := ic.current
ic.current = nil
return signal
}
// Peek returns the current interrupt without removing it
func (ic *InterruptController) Peek() *InterruptSignal {
if ic == nil {
return nil
}
ic.mutex.RLock()
defer ic.mutex.RUnlock()
return ic.current
}
// IsInterrupted checks if interrupt context is cancelled (force interrupt)
func (ic *InterruptController) IsInterrupted() bool {
if ic == nil || ic.ctx == nil {
return false
}
select {
case <-ic.ctx.Done():
return true
default:
return false
}
}
// Context returns the interrupt control context
// This can be used in select statements to check for force interrupts
func (ic *InterruptController) Context() context.Context {
if ic == nil {
return context.Background()
}
return ic.ctx
}
// GetPendingCount returns the number of pending interrupts
func (ic *InterruptController) GetPendingCount() int {
if ic == nil {
return 0
}
ic.mutex.RLock()
defer ic.mutex.RUnlock()
count := len(ic.pending)
if ic.current != nil {
count++
}
return count
}
// Clear clears all interrupts (current and pending)
func (ic *InterruptController) Clear() {
if ic == nil {
return
}
ic.mutex.Lock()
defer ic.mutex.Unlock()
ic.current = nil
ic.pending = make([]*InterruptSignal, 0)
}
// Stop stops the interrupt controller and cleans up resources
func (ic *InterruptController) Stop() {
if ic == nil {
return
}
// Cancel context to stop listener
if ic.cancel != nil {
ic.cancel()
}
// Close channel
if ic.queue != nil {
close(ic.queue)
}
// Clear interrupts
ic.Clear()
}
// SendSignal sends an interrupt signal to the controller
// This is called from external sources (e.g., another HTTP request)
func (ic *InterruptController) SendSignal(signal *InterruptSignal) error {
if ic == nil {
return fmt.Errorf("interrupt controller is nil")
}
if ic.queue == nil {
return fmt.Errorf("interrupt queue is not initialized")
}
// Non-blocking send
select {
case ic.queue <- signal:
return nil
case <-time.After(500 * time.Millisecond):
return fmt.Errorf("failed to send interrupt: timeout")
}
}

View file

@ -0,0 +1,661 @@
package context
import (
stdContext "context"
"fmt"
"testing"
"time"
"github.com/yaoapp/gou/plan"
"github.com/yaoapp/yao/openapi/oauth/types"
)
// newTestContextWithInterrupt creates a Context with interrupt controller for testing
func newTestContextWithInterrupt(chatID, assistantID string) *Context {
ctx := &Context{
Context: stdContext.Background(),
ID: fmt.Sprintf("test_ctx_%d", time.Now().UnixNano()),
Space: plan.NewMemorySharedSpace(),
ChatID: chatID,
AssistantID: assistantID,
Connector: "",
Locale: "en-us",
Theme: "light",
Client: Client{
Type: "web",
UserAgent: "TestAgent/1.0",
IP: "127.0.0.1",
},
Referer: RefererAPI,
Accept: AcceptWebCUI,
Route: "/test/route",
Metadata: map[string]interface{}{
"test": "context_metadata",
},
Authorized: &types.AuthorizedInfo{
Subject: "test-user",
ClientID: "test-client-id",
UserID: "test-user-123",
TeamID: "test-team-456",
TenantID: "test-tenant-789",
SessionID: "test-session-id",
},
}
// Initialize interrupt controller
ctx.Interrupt = NewInterruptController()
ctx.Interrupt.SetContextID(ctx.ID)
// Register context globally
if err := Register(ctx); err != nil {
panic(fmt.Sprintf("Failed to register context: %v", err))
}
// Start interrupt listener
ctx.Interrupt.Start()
return ctx
}
// TestInterruptBasic tests basic interrupt signal sending and receiving
func TestInterruptBasic(t *testing.T) {
// Create context with interrupt support
ctx := newTestContextWithInterrupt("chat-test-interrupt", "test-assistant")
defer ctx.Release()
t.Run("SendGracefulInterrupt", func(t *testing.T) {
// Create a graceful interrupt signal
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{
{Role: RoleUser, Content: "This is a graceful interrupt"},
},
Timestamp: time.Now().UnixMilli(),
}
// Send interrupt signal
err := SendInterrupt(ctx.ID, signal)
if err != nil {
t.Fatalf("Failed to send interrupt signal: %v", err)
}
// Wait a bit for the signal to be processed
time.Sleep(100 * time.Millisecond)
// Check if signal was received
receivedSignal := ctx.Interrupt.Peek()
if receivedSignal == nil {
t.Fatal("Expected interrupt signal to be received, got nil")
}
if receivedSignal.Type != InterruptGraceful {
t.Errorf("Expected interrupt type 'graceful', got: %s", receivedSignal.Type)
}
if len(receivedSignal.Messages) != 1 {
t.Errorf("Expected 1 message, got: %d", len(receivedSignal.Messages))
}
if receivedSignal.Messages[0].Content != "This is a graceful interrupt" {
t.Errorf("Expected message content 'This is a graceful interrupt', got: %s", receivedSignal.Messages[0].Content)
}
t.Log("✓ Graceful interrupt signal sent and received successfully")
})
t.Run("SendForceInterrupt", func(t *testing.T) {
// Clear previous signals
ctx.Interrupt.Clear()
// Create a force interrupt signal
signal := &InterruptSignal{
Type: InterruptForce,
Messages: []Message{
{Role: RoleUser, Content: "This is a force interrupt"},
},
Timestamp: time.Now().UnixMilli(),
}
// Send interrupt signal
err := SendInterrupt(ctx.ID, signal)
if err != nil {
t.Fatalf("Failed to send interrupt signal: %v", err)
}
// Wait a bit for the signal to be processed
time.Sleep(100 * time.Millisecond)
// Check if signal was received
receivedSignal := ctx.Interrupt.Peek()
if receivedSignal == nil {
t.Fatal("Expected interrupt signal to be received, got nil")
}
if receivedSignal.Type != InterruptForce {
t.Errorf("Expected interrupt type 'force', got: %s", receivedSignal.Type)
}
t.Log("✓ Force interrupt signal sent and received successfully")
})
t.Run("MultipleInterrupts", func(t *testing.T) {
// Clear previous signals
ctx.Interrupt.Clear()
// Send multiple interrupt signals
for i := 0; i < 3; i++ {
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{
{Role: RoleUser, Content: fmt.Sprintf("Message %d", i+1)},
},
Timestamp: time.Now().UnixMilli(),
}
err := SendInterrupt(ctx.ID, signal)
if err != nil {
t.Fatalf("Failed to send interrupt signal %d: %v", i+1, err)
}
}
// Wait a bit for signals to be processed
time.Sleep(100 * time.Millisecond)
// Check pending count
pendingCount := ctx.Interrupt.GetPendingCount()
if pendingCount != 3 {
t.Errorf("Expected 3 pending interrupts, got: %d", pendingCount)
}
// Check merged signal
mergedSignal := ctx.Interrupt.CheckWithMerge()
if mergedSignal == nil {
t.Fatal("Expected merged signal, got nil")
}
if len(mergedSignal.Messages) != 3 {
t.Errorf("Expected 3 merged messages, got: %d", len(mergedSignal.Messages))
}
// Verify all messages are present
for i := 0; i < 3; i++ {
expectedContent := fmt.Sprintf("Message %d", i+1)
if mergedSignal.Messages[i].Content != expectedContent {
t.Errorf("Expected message %d content '%s', got: %s", i+1, expectedContent, mergedSignal.Messages[i].Content)
}
}
t.Log("✓ Multiple interrupt signals merged successfully")
})
}
// TestInterruptHandler tests interrupt handler invocation
func TestInterruptHandler(t *testing.T) {
// Create context with interrupt support
ctx := newTestContextWithInterrupt("chat-test-interrupt-handler", "test-assistant")
defer ctx.Release()
t.Run("HandlerInvocation", func(t *testing.T) {
// Track if handler was called
handlerCalled := false
var receivedSignal *InterruptSignal
// Set up handler
ctx.Interrupt.SetHandler(func(c *Context, signal *InterruptSignal) error {
handlerCalled = true
receivedSignal = signal
t.Logf("Handler called with signal type: %s, messages: %d", signal.Type, len(signal.Messages))
return nil
})
// Send interrupt signal
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{
{Role: RoleUser, Content: "Test handler invocation"},
},
Timestamp: time.Now().UnixMilli(),
}
err := SendInterrupt(ctx.ID, signal)
if err != nil {
t.Fatalf("Failed to send interrupt signal: %v", err)
}
// Wait for handler to be called
time.Sleep(200 * time.Millisecond)
// Verify handler was called
if !handlerCalled {
t.Error("Expected handler to be called, but it wasn't")
}
if receivedSignal == nil {
t.Fatal("Expected signal in handler, got nil")
}
if receivedSignal.Type != InterruptGraceful {
t.Errorf("Expected graceful interrupt in handler, got: %s", receivedSignal.Type)
}
if len(receivedSignal.Messages) != 1 {
t.Errorf("Expected 1 message in handler, got: %d", len(receivedSignal.Messages))
}
t.Log("✓ Interrupt handler invoked successfully")
})
t.Run("HandlerWithError", func(t *testing.T) {
// Create new context
ctx2 := newTestContextWithInterrupt("chat-test-handler-error", "test-assistant")
defer ctx2.Release()
// Set up handler that returns error
handlerCalled := false
ctx2.Interrupt.SetHandler(func(c *Context, signal *InterruptSignal) error {
handlerCalled = true
return fmt.Errorf("test error from handler")
})
// Send interrupt signal
signal := &InterruptSignal{
Type: InterruptForce,
Messages: []Message{
{Role: RoleUser, Content: "Test error handling"},
},
Timestamp: time.Now().UnixMilli(),
}
err := SendInterrupt(ctx2.ID, signal)
if err != nil {
t.Fatalf("Failed to send interrupt signal: %v", err)
}
// Wait for handler to be called
time.Sleep(200 * time.Millisecond)
// Handler should still be called even if it returns error
if !handlerCalled {
t.Error("Expected handler to be called even with error")
}
t.Log("✓ Handler error handling works correctly")
})
}
// TestInterruptContextLifecycle tests context registration and cleanup
func TestInterruptContextLifecycle(t *testing.T) {
t.Run("RegisterAndRetrieve", func(t *testing.T) {
ctx := newTestContextWithInterrupt("chat-test-lifecycle", "test-assistant")
// Verify context can be retrieved
retrievedCtx, err := Get(ctx.ID)
if err != nil {
t.Fatalf("Failed to retrieve context: %v", err)
}
if retrievedCtx.ID != ctx.ID {
t.Errorf("Expected context ID %s, got: %s", ctx.ID, retrievedCtx.ID)
}
ctx.Release()
// After release, context should be removed
_, err = Get(ctx.ID)
if err == nil {
t.Error("Expected error when retrieving released context")
}
t.Log("✓ Context registration and cleanup works correctly")
})
t.Run("SendToNonExistentContext", func(t *testing.T) {
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{{Role: RoleUser, Content: "test"}},
Timestamp: time.Now().UnixMilli(),
}
err := SendInterrupt("non-existent-id", signal)
if err == nil {
t.Error("Expected error when sending to non-existent context")
}
t.Log("✓ Sending to non-existent context returns error")
})
}
// TestInterruptCheckMethods tests different check methods
func TestInterruptCheckMethods(t *testing.T) {
ctx := newTestContextWithInterrupt("chat-test-check-methods", "test-assistant")
defer ctx.Release()
t.Run("PeekDoesNotRemove", func(t *testing.T) {
// Send signal
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{{Role: RoleUser, Content: "peek test"}},
Timestamp: time.Now().UnixMilli(),
}
SendInterrupt(ctx.ID, signal)
time.Sleep(100 * time.Millisecond)
// Peek should return signal but not remove it
peeked1 := ctx.Interrupt.Peek()
if peeked1 == nil {
t.Fatal("Expected signal from first peek")
}
peeked2 := ctx.Interrupt.Peek()
if peeked2 == nil {
t.Fatal("Expected signal from second peek")
}
if peeked1.Messages[0].Content != peeked2.Messages[0].Content {
t.Error("Peek should return the same signal")
}
t.Log("✓ Peek does not remove signal")
})
t.Run("CheckRemovesSignal", func(t *testing.T) {
ctx.Interrupt.Clear()
// Send signal
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{{Role: RoleUser, Content: "check test"}},
Timestamp: time.Now().UnixMilli(),
}
SendInterrupt(ctx.ID, signal)
time.Sleep(100 * time.Millisecond)
// Check should return and remove signal
checked := ctx.Interrupt.Check()
if checked == nil {
t.Fatal("Expected signal from check")
}
// Second check should return nil
checked2 := ctx.Interrupt.Check()
if checked2 != nil {
t.Error("Expected nil from second check after removal")
}
t.Log("✓ Check removes signal after retrieval")
})
t.Run("CheckWithMergeMultipleSignals", func(t *testing.T) {
ctx.Interrupt.Clear()
// Send 5 signals with different messages
messages := []string{
"First message",
"Second message",
"Third message",
"Fourth message",
"Fifth message",
}
for i, msg := range messages {
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{
{Role: RoleUser, Content: msg},
},
Timestamp: time.Now().UnixMilli(),
Metadata: map[string]interface{}{
"sequence": i + 1,
},
}
err := SendInterrupt(ctx.ID, signal)
if err != nil {
t.Fatalf("Failed to send signal %d: %v", i+1, err)
}
time.Sleep(10 * time.Millisecond) // Small delay between signals
}
time.Sleep(100 * time.Millisecond)
// Verify all signals are queued
pendingCount := ctx.Interrupt.GetPendingCount()
if pendingCount != 5 {
t.Errorf("Expected 5 pending signals, got: %d", pendingCount)
}
// CheckWithMerge should merge all messages into one signal
merged := ctx.Interrupt.CheckWithMerge()
if merged == nil {
t.Fatal("Expected merged signal, got nil")
}
// Verify all messages are merged
if len(merged.Messages) != 5 {
t.Errorf("Expected 5 merged messages, got: %d", len(merged.Messages))
}
// Verify message order
for i, msg := range messages {
if merged.Messages[i].Content != msg {
t.Errorf("Message %d mismatch: expected '%s', got '%s'", i+1, msg, merged.Messages[i].Content)
}
}
// Verify metadata indicates merge
if merged.Metadata["merged"] != true {
t.Error("Expected merged metadata to be true")
}
if merged.Metadata["merged_count"] != 5 {
t.Errorf("Expected merged_count 5, got: %v", merged.Metadata["merged_count"])
}
// After merge, queue should be empty
if ctx.Interrupt.GetPendingCount() != 0 {
t.Errorf("Expected empty queue after merge, got: %d", ctx.Interrupt.GetPendingCount())
}
t.Log("✓ CheckWithMerge correctly merged 5 signals into one")
})
t.Run("CheckWithMergeSingleSignal", func(t *testing.T) {
ctx.Interrupt.Clear()
// Send single signal
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{{Role: RoleUser, Content: "single signal"}},
Timestamp: time.Now().UnixMilli(),
}
SendInterrupt(ctx.ID, signal)
time.Sleep(100 * time.Millisecond)
// CheckWithMerge with single signal should return it without merge metadata
merged := ctx.Interrupt.CheckWithMerge()
if merged == nil {
t.Fatal("Expected signal, got nil")
}
if len(merged.Messages) != 1 {
t.Errorf("Expected 1 message, got: %d", len(merged.Messages))
}
// Single signal should not have merge metadata
if merged.Metadata != nil && merged.Metadata["merged"] == true {
t.Error("Single signal should not have merge metadata")
}
t.Log("✓ CheckWithMerge handles single signal correctly")
})
}
// TestInterruptContext tests interrupt context methods
func TestInterruptContext(t *testing.T) {
ctx := newTestContextWithInterrupt("chat-test-interrupt-context", "test-assistant")
defer ctx.Release()
t.Run("InterruptContextMethod", func(t *testing.T) {
// Get interrupt context
interruptCtx := ctx.Interrupt.Context()
if interruptCtx == nil {
t.Fatal("Expected interrupt context, got nil")
}
// Context should not be done initially
select {
case <-interruptCtx.Done():
t.Error("Interrupt context should not be done initially")
default:
t.Log("✓ Interrupt context is not done initially")
}
})
t.Run("IsInterruptedFalseInitially", func(t *testing.T) {
// Should not be interrupted initially
if ctx.Interrupt.IsInterrupted() {
t.Error("Should not be interrupted initially")
}
t.Log("✓ IsInterrupted returns false initially")
})
t.Run("ForceInterruptCancelsContext", func(t *testing.T) {
// Get context before interrupt
interruptCtx := ctx.Interrupt.Context()
// Send force interrupt
signal := &InterruptSignal{
Type: InterruptForce,
Messages: []Message{{Role: RoleUser, Content: "force stop"}},
Timestamp: time.Now().UnixMilli(),
}
err := SendInterrupt(ctx.ID, signal)
if err != nil {
t.Fatalf("Failed to send interrupt: %v", err)
}
time.Sleep(100 * time.Millisecond)
// The OLD context should be cancelled
select {
case <-interruptCtx.Done():
t.Log("✓ Force interrupt cancelled the old context")
case <-time.After(200 * time.Millisecond):
t.Error("Old context was not cancelled after force interrupt")
}
// Note: IsInterrupted() checks the NEW context (which was recreated)
// So it will return false. This is expected behavior.
// The key is that the old context was cancelled (checked above)
t.Log("✓ Context was recreated after force interrupt (expected behavior)")
})
t.Run("GracefulInterruptDoesNotCancelContext", func(t *testing.T) {
// Create new context for this test
ctx2 := newTestContextWithInterrupt("chat-test-graceful-no-cancel", "test-assistant")
defer ctx2.Release()
interruptCtx := ctx2.Interrupt.Context()
// Send graceful interrupt
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{{Role: RoleUser, Content: "graceful"}},
Timestamp: time.Now().UnixMilli(),
}
SendInterrupt(ctx2.ID, signal)
time.Sleep(100 * time.Millisecond)
// Context should NOT be cancelled for graceful interrupt
select {
case <-interruptCtx.Done():
t.Error("Graceful interrupt should not cancel context")
default:
t.Log("✓ Graceful interrupt does not cancel context")
}
// IsInterrupted should still return false for graceful
if ctx2.Interrupt.IsInterrupted() {
t.Error("IsInterrupted should return false for graceful interrupt")
} else {
t.Log("✓ IsInterrupted returns false for graceful interrupt")
}
})
}
// TestInterruptSendSignalDirectly tests SendSignal method directly
func TestInterruptSendSignalDirectly(t *testing.T) {
ctx := newTestContextWithInterrupt("chat-test-send-signal", "test-assistant")
defer ctx.Release()
t.Run("SendSignalSuccess", func(t *testing.T) {
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{{Role: RoleUser, Content: "direct send"}},
Timestamp: time.Now().UnixMilli(),
}
err := ctx.Interrupt.SendSignal(signal)
if err != nil {
t.Fatalf("SendSignal failed: %v", err)
}
time.Sleep(100 * time.Millisecond)
// Verify signal was received
received := ctx.Interrupt.Peek()
if received == nil {
t.Fatal("Signal not received")
}
if received.Messages[0].Content != "direct send" {
t.Errorf("Expected 'direct send', got: %s", received.Messages[0].Content)
}
t.Log("✓ SendSignal directly works")
})
t.Run("SendSignalToNilController", func(t *testing.T) {
var nilController *InterruptController
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{{Role: RoleUser, Content: "test"}},
Timestamp: time.Now().UnixMilli(),
}
err := nilController.SendSignal(signal)
if err == nil {
t.Error("Expected error when sending to nil controller")
} else {
t.Logf("✓ Correctly returned error for nil controller: %v", err)
}
})
t.Run("SendSignalTimeout", func(t *testing.T) {
// Create controller but don't start listener
testCtrl := NewInterruptController()
// Don't call Start(), so channel won't be read
// Fill the buffer (capacity is 10)
for i := 0; i < 10; i++ {
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{{Role: RoleUser, Content: fmt.Sprintf("msg %d", i)}},
Timestamp: time.Now().UnixMilli(),
}
testCtrl.SendSignal(signal)
}
// This should timeout since buffer is full and no listener
signal := &InterruptSignal{
Type: InterruptGraceful,
Messages: []Message{{Role: RoleUser, Content: "overflow"}},
Timestamp: time.Now().UnixMilli(),
}
err := testCtrl.SendSignal(signal)
if err == nil {
t.Error("Expected timeout error when buffer is full")
} else {
t.Logf("✓ SendSignal correctly times out when buffer full: %v", err)
}
})
}

View file

@ -44,9 +44,10 @@ func GetCompletionRequest(c *gin.Context, cache store.Store) (*CompletionRequest
clientType := getClientType(userAgent)
clientIP := c.ClientIP()
// Set cache in context
// Create context with unique ID
ctx := &Context{
Context: c.Request.Context(),
ID: generateContextID(),
Space: plan.NewMemorySharedSpace(),
Cache: cache,
Writer: c.Writer,
@ -66,6 +67,18 @@ func GetCompletionRequest(c *gin.Context, cache store.Store) (*CompletionRequest
Metadata: GetMetadata(c, completionReq),
}
// Initialize interrupt controller
ctx.Interrupt = NewInterruptController()
ctx.Interrupt.SetContextID(ctx.ID)
// Register context to global registry first
if err := Register(ctx); err != nil {
return nil, nil, fmt.Errorf("failed to register context: %w", err)
}
// Start interrupt listener after registration
ctx.Interrupt.Start()
return completionReq, ctx, nil
}

View file

@ -2,6 +2,7 @@ package context
import (
"context"
"sync"
"github.com/yaoapp/gou/plan"
"github.com/yaoapp/gou/store"
@ -116,17 +117,81 @@ var ValidStackStatus = map[string]bool{
StackStatusTimeout: true,
}
// Interrupt Types and Constants
// ===============================
// InterruptType represents the type of interrupt
type InterruptType string
const (
// InterruptGraceful waits for current step to complete before handling interrupt
InterruptGraceful InterruptType = "graceful"
// InterruptForce immediately cancels current operation and handles interrupt
InterruptForce InterruptType = "force"
)
// InterruptAction represents the action to take after interrupt is handled
type InterruptAction string
const (
// InterruptActionContinue appends new messages and continues execution
InterruptActionContinue InterruptAction = "continue"
// InterruptActionRestart restarts execution with only new messages
InterruptActionRestart InterruptAction = "restart"
// InterruptActionAbort terminates the request
InterruptActionAbort InterruptAction = "abort"
)
// InterruptSignal represents an interrupt signal with new messages from user
type InterruptSignal struct {
Type InterruptType `json:"type"` // Interrupt type: graceful or force
Messages []Message `json:"messages"` // User's new messages (can be multiple)
Timestamp int64 `json:"timestamp"` // Interrupt timestamp in milliseconds
Metadata map[string]interface{} `json:"metadata,omitempty"` // Additional metadata
}
// InterruptHandler is the function signature for handling interrupts
// This handler is registered in the InterruptController and called when interrupt signal is received
// Parameters:
// - ctx: The context being interrupted
// - signal: The interrupt signal (contains Type and Messages)
//
// Returns:
// - error: Error if interrupt handling failed
type InterruptHandler func(ctx *Context, signal *InterruptSignal) error
// InterruptController manages interrupt handling for a context
// All interrupt-related fields are encapsulated in this type
type InterruptController struct {
queue chan *InterruptSignal `json:"-"` // Queue to receive interrupt signals
current *InterruptSignal `json:"-"` // Current interrupt being processed
pending []*InterruptSignal `json:"-"` // Pending interrupts in queue
mutex sync.RWMutex `json:"-"` // Protects current and pending
ctx context.Context `json:"-"` // Interrupt control context (independent from HTTP context)
cancel context.CancelFunc `json:"-"` // Cancel function for force interrupt
listenerStarted bool `json:"-"` // Whether listener goroutine is started
handler InterruptHandler `json:"-"` // Handler to process interrupt signals
contextID string `json:"-"` // Context ID to retrieve the parent context
}
// Context the context
type Context struct {
// Context
context.Context
Space plan.Space `json:"-"` // Shared data space, it will be used to share data between the request and the call
Cache store.Store `json:"-"` // Cache store, it will be used to store the message cache, default is "__yao.agent.cache"
Stack *Stack `json:"-"` // Stack, current active stack of the request
Stacks map[string]*Stack `json:"-"` // Stacks, all stacks in this request (for trace logging)
Writer Writer `json:"-"` // Writer, it will be used to write response data to the client
trace traceTypes.Manager `json:"-"` // Trace manager, lazy initialized on first access
ID string `json:"id"` // Context ID for external interrupt identification
Space plan.Space `json:"-"` // Shared data space, it will be used to share data between the request and the call
Cache store.Store `json:"-"` // Cache store, it will be used to store the message cache, default is "__yao.agent.cache"
Stack *Stack `json:"-"` // Stack, current active stack of the request
Stacks map[string]*Stack `json:"-"` // Stacks, all stacks in this request (for trace logging)
Writer Writer `json:"-"` // Writer, it will be used to write response data to the client
trace traceTypes.Manager `json:"-"` // Trace manager, lazy initialized on first access
// Interrupt control (all interrupt-related logic is encapsulated in InterruptController)
Interrupt *InterruptController `json:"-"` // Interrupt controller for handling user interrupts during streaming
// Authorized information
Authorized *types.AuthorizedInfo `json:"authorized,omitempty"` // Authorized information
@ -225,6 +290,26 @@ type ResponseHookMCP struct{}
// ResponseHookFailback the response of the failback hook
type ResponseHookFailback struct{}
// HookInterruptedResponse the response of the interrupted hook
type HookInterruptedResponse struct {
// Action to take after interrupt is handled
Action InterruptAction `json:"action"` // continue, restart, or abort
// Messages to use for next execution (if action is continue or restart)
Messages []Message `json:"messages,omitempty"`
// Context adjustments - allow hook to modify context fields
AssistantID string `json:"assistant_id,omitempty"` // Override assistant ID
Connector string `json:"connector,omitempty"` // Override connector
Locale string `json:"locale,omitempty"` // Override locale
Theme string `json:"theme,omitempty"` // Override theme
Route string `json:"route,omitempty"` // Override route
Metadata map[string]interface{} `json:"metadata,omitempty"` // Override or merge metadata
// Notice to send to client
Notice string `json:"notice,omitempty"` // Message to display to user (e.g., "Processing your new question...")
}
// Message Structure ( OpenAI Chat Completion Input Message Structure, https://platform.openai.com/docs/api-reference/chat/create#chat/create-messages )
// ===============================

View file

@ -209,6 +209,13 @@ func (p *Provider) Stream(ctx *context.Context, messages []context.Message, opti
default:
}
// Check for force interrupt before retry
if ctx.Interrupt != nil {
if signal := ctx.Interrupt.Peek(); signal != nil && signal.Type == context.InterruptForce {
return nil, fmt.Errorf("force interrupted by user")
}
}
if attempt > 0 {
// Exponential backoff: 1s, 2s, 4s
backoff := time.Duration(1<<uint(attempt-1)) * time.Second
@ -216,12 +223,27 @@ func (p *Provider) Stream(ctx *context.Context, messages []context.Message, opti
// Sleep with context cancellation support
timer := time.NewTimer(backoff)
select {
case <-timer.C:
// Continue to retry
case <-goCtx.Done():
timer.Stop()
return nil, fmt.Errorf("context cancelled during backoff: %w", goCtx.Err())
interruptTicker := time.NewTicker(100 * time.Millisecond) // Check interrupt every 100ms
defer interruptTicker.Stop()
backoffLoop:
for {
select {
case <-timer.C:
// Backoff completed, continue to retry
break backoffLoop
case <-goCtx.Done():
timer.Stop()
return nil, fmt.Errorf("context cancelled during backoff: %w", goCtx.Err())
case <-interruptTicker.C:
// Check for force interrupt during backoff
if ctx.Interrupt != nil {
if signal := ctx.Interrupt.Peek(); signal != nil && signal.Type == context.InterruptForce {
timer.Stop()
return nil, fmt.Errorf("force interrupted by user during backoff")
}
}
}
}
}
@ -292,6 +314,13 @@ func (p *Provider) streamWithRetry(ctx *context.Context, messages []context.Mess
default:
}
// Check for force interrupt before stream start
if ctx.Interrupt != nil {
if signal := ctx.Interrupt.Peek(); signal != nil && signal.Type == context.InterruptForce {
return nil, fmt.Errorf("force interrupted by user before stream start")
}
}
// Send stream_start event
if handler != nil {
model, _ := p.GetModel()
@ -389,6 +418,14 @@ func (p *Provider) streamWithRetry(ctx *context.Context, messages []context.Mess
default:
}
// Check for force interrupt signal
if ctx.Interrupt != nil {
if signal := ctx.Interrupt.Peek(); signal != nil && signal.Type == context.InterruptForce {
log.Warn("Stream cancelled by force interrupt")
return http.HandlerReturnBreak
}
}
if len(data) == 0 {
return http.HandlerReturnOk
}