From c0e7ddbf49174826a125a50509ae4d55662e7e94 Mon Sep 17 00:00:00 2001 From: Lakshman Patel Date: Thu, 20 Aug 2026 23:16:58 +0530 Subject: [PATCH 1/2] feat(engine): add prompt queue, announcements, and transcript resume Port of Year 0 PACK-02 and PACK-08 control plane capabilities into Hawk. - engine: Implemented thread-safe PromptQueue supporting priority scheduling (Normal, Steering, Interjection) and FIFO ordering. - engine: Implemented AnnouncementFeed managing in-session broadcast notices with TTL expiration. - engine: Wired PromptQueue and AnnouncementFeed accessors to Session. - engine: Implemented true transcript resume in agent_session_tool.go by replaying prior session messages from disk when ResumeFrom is provided. - tool: Added unit tests for structured AskUserQuestionTool option handling and validation. - tests: Added test suites for prompt queue priority/concurrency, announcement feed expiration, and subagent transcript resume. --- internal/engine/agent_resume_test.go | 109 ++++++++++++++ internal/engine/agent_session_tool.go | 10 +- internal/engine/announcements.go | 134 +++++++++++++++++ internal/engine/announcements_test.go | 63 ++++++++ internal/engine/persistence_service.go | 17 +++ internal/engine/prompt_queue.go | 190 +++++++++++++++++++++++++ internal/engine/prompt_queue_test.go | 154 ++++++++++++++++++++ internal/engine/session.go | 28 ++++ internal/tool/ask_user_test.go | 98 +++++++++++++ 9 files changed, 801 insertions(+), 2 deletions(-) create mode 100644 internal/engine/agent_resume_test.go create mode 100644 internal/engine/announcements.go create mode 100644 internal/engine/announcements_test.go create mode 100644 internal/engine/prompt_queue.go create mode 100644 internal/engine/prompt_queue_test.go create mode 100644 internal/tool/ask_user_test.go diff --git a/internal/engine/agent_resume_test.go b/internal/engine/agent_resume_test.go new file mode 100644 index 00000000..cb9fe860 --- /dev/null +++ b/internal/engine/agent_resume_test.go @@ -0,0 +1,109 @@ +package engine + +import ( + "testing" + + agentcontracts "github.com/GrayCodeAI/hawk-core-contracts/agent" + "github.com/GrayCodeAI/hawk/internal/session" + "github.com/GrayCodeAI/hawk/internal/storage" + "github.com/GrayCodeAI/hawk/internal/tool" +) + +func TestSubAgentResume_ReplaysTranscriptMessages(t *testing.T) { + tempDir := t.TempDir() + storage.SetTestDirs(t, tempDir) + + // 1. Create and save a prior session + prior := &session.Session{ + ID: "subagent-prior-123", + Model: "test-model", + Messages: []session.Message{ + {Role: "user", Content: "Find all auth handlers."}, + {Role: "assistant", Content: "Auth handlers are in internal/auth/handler.go."}, + }, + } + if err := session.Save(prior); err != nil { + t.Fatalf("failed to save prior session: %v", err) + } + + // 2. Set up parent session + reg := tool.NewRegistry() + parent := NewSession("", "", "You are parent assistant", reg) + + req := agentcontracts.SpawnRequest{ + Prompt: "Where is the token validation function?", + SubagentType: "explore", + ResumeFrom: "subagent-prior-123", + } + + norm, err := req.Normalize() + if err != nil { + t.Fatalf("Normalize failed: %v", err) + } + + sub := parent.SubSession("", "", reg) + if norm.ResumeFrom != "" { + if priorSession, loadErr := session.Load(norm.ResumeFrom); loadErr == nil && priorSession != nil { + for _, m := range priorSession.Messages { + sub.Persistence().AddMessage(m.Role, m.Content) + } + } + } + sub.AddUser(norm.Prompt) + + // 3. Verify that the sub-session transcript has the restored prior messages + msgs := sub.Persistence().Messages() + if len(msgs) != 3 { + t.Fatalf("expected 3 messages in sub-session transcript, got %d", len(msgs)) + } + + if msgs[0].Role != "user" || msgs[0].Content != "Find all auth handlers." { + t.Errorf("msg[0] = %+v, want user 'Find all auth handlers.'", msgs[0]) + } + if msgs[1].Role != "assistant" || msgs[1].Content != "Auth handlers are in internal/auth/handler.go." { + t.Errorf("msg[1] = %+v, want assistant findings", msgs[1]) + } + if msgs[2].Role != "user" || msgs[2].Content != "Where is the token validation function?" { + t.Errorf("msg[2] = %+v, want user new prompt", msgs[2]) + } +} + +func TestSubAgentResume_FallbackOnMissingSession(t *testing.T) { + tempDir := t.TempDir() + storage.SetTestDirs(t, tempDir) + + reg := tool.NewRegistry() + parent := NewSession("", "", "You are parent assistant", reg) + + req := agentcontracts.SpawnRequest{ + Prompt: "Continue analysis.", + SubagentType: "explore", + ResumeFrom: "nonexistent-subagent-999", + } + + norm, err := req.Normalize() + if err != nil { + t.Fatalf("Normalize failed: %v", err) + } + + sub := parent.SubSession("", "", reg) + prompt := norm.Prompt + if norm.ResumeFrom != "" { + if priorSession, loadErr := session.Load(norm.ResumeFrom); loadErr == nil && priorSession != nil { + for _, m := range priorSession.Messages { + sub.Persistence().AddMessage(m.Role, m.Content) + } + } else { + prompt = "Resume prior subagent " + norm.ResumeFrom + ".\n\n" + prompt + } + } + sub.AddUser(prompt) + + msgs := sub.Persistence().Messages() + if len(msgs) != 1 { + t.Fatalf("expected 1 message in fallback, got %d", len(msgs)) + } + if msgs[0].Content != "Resume prior subagent nonexistent-subagent-999.\n\nContinue analysis." { + t.Errorf("unexpected fallback prompt: %q", msgs[0].Content) + } +} diff --git a/internal/engine/agent_session_tool.go b/internal/engine/agent_session_tool.go index c6889f16..481b7a6d 100644 --- a/internal/engine/agent_session_tool.go +++ b/internal/engine/agent_session_tool.go @@ -15,6 +15,7 @@ import ( "github.com/GrayCodeAI/hawk/internal/hooks" "github.com/GrayCodeAI/hawk/internal/prompts" "github.com/GrayCodeAI/hawk/internal/sandbox" + "github.com/GrayCodeAI/hawk/internal/session" "github.com/GrayCodeAI/hawk/internal/tool" ) @@ -247,8 +248,13 @@ func (s *Session) spawnSubAgent(ctx context.Context, norm agentcontracts.Normali prompt = fmt.Sprintf("Working directory: %s\n\n%s", workDir, prompt) } if norm.ResumeFrom != "" { - // True transcript resume lands with taskruntime persistence; surface the id. - prompt = fmt.Sprintf("Resume prior subagent %s.\n\n%s", norm.ResumeFrom, prompt) + if priorSession, loadErr := session.Load(norm.ResumeFrom); loadErr == nil && priorSession != nil { + for _, m := range priorSession.Messages { + sub.Persistence().AddMessage(m.Role, m.Content) + } + } else { + prompt = fmt.Sprintf("Resume prior subagent %s.\n\n%s", norm.ResumeFrom, prompt) + } } sub.AddUser(prompt) diff --git a/internal/engine/announcements.go b/internal/engine/announcements.go new file mode 100644 index 00000000..1e9be2c2 --- /dev/null +++ b/internal/engine/announcements.go @@ -0,0 +1,134 @@ +package engine + +import ( + "crypto/rand" + "encoding/hex" + "sync" + "time" +) + +// AnnouncementKind denotes the category/urgency of an in-session announcement. +type AnnouncementKind string + +const ( + AnnouncementInfo AnnouncementKind = "info" + AnnouncementWarning AnnouncementKind = "warning" + AnnouncementSystem AnnouncementKind = "system" + AnnouncementSchedule AnnouncementKind = "schedule" +) + +// Announcement represents a single broadcast notice within an active session. +type Announcement struct { + ID string `json:"id"` + Kind AnnouncementKind `json:"kind"` + Message string `json:"message"` + CreatedAt time.Time `json:"created_at"` + ExpiresAt time.Time `json:"expires_at,omitempty"` + Read bool `json:"read"` +} + +// AnnouncementFeed provides a thread-safe registry of active session announcements. +type AnnouncementFeed struct { + mu sync.RWMutex + announcements []*Announcement +} + +// NewAnnouncementFeed creates an empty AnnouncementFeed. +func NewAnnouncementFeed() *AnnouncementFeed { + return &AnnouncementFeed{ + announcements: make([]*Announcement, 0), + } +} + +// Post broadcasts a new announcement with an optional TTL (0 = never expires). +func (af *AnnouncementFeed) Post(kind AnnouncementKind, message string, ttl time.Duration) *Announcement { + af.mu.Lock() + defer af.mu.Unlock() + + now := time.Now() + var expiresAt time.Time + if ttl > 0 { + expiresAt = now.Add(ttl) + } + + a := &Announcement{ + ID: generateAnnouncementID(), + Kind: kind, + Message: message, + CreatedAt: now, + ExpiresAt: expiresAt, + Read: false, + } + + af.announcements = append(af.announcements, a) + return a +} + +// Active returns all non-expired announcements. +func (af *AnnouncementFeed) Active() []Announcement { + af.mu.Lock() + defer af.mu.Unlock() + + now := time.Now() + active := make([]Announcement, 0, len(af.announcements)) + remaining := make([]*Announcement, 0, len(af.announcements)) + + for _, a := range af.announcements { + if a.ExpiresAt.IsZero() || a.ExpiresAt.After(now) { + active = append(active, *a) + remaining = append(remaining, a) + } + } + + af.announcements = remaining + return active +} + +// Unread returns all active announcements that have not been acknowledged. +func (af *AnnouncementFeed) Unread() []Announcement { + active := af.Active() + unread := make([]Announcement, 0, len(active)) + for _, a := range active { + if !a.Read { + unread = append(unread, a) + } + } + return unread +} + +// MarkRead marks an announcement as read by its ID. +func (af *AnnouncementFeed) MarkRead(id string) bool { + af.mu.Lock() + defer af.mu.Unlock() + + for _, a := range af.announcements { + if a.ID == id { + a.Read = true + return true + } + } + return false +} + +// MarkAllRead marks all active announcements as read. +func (af *AnnouncementFeed) MarkAllRead() { + af.mu.Lock() + defer af.mu.Unlock() + + for _, a := range af.announcements { + a.Read = true + } +} + +// Clear removes all announcements. +func (af *AnnouncementFeed) Clear() { + af.mu.Lock() + defer af.mu.Unlock() + af.announcements = af.announcements[:0] +} + +func generateAnnouncementID() string { + b := make([]byte, 6) + _, _ = rand.Read(b) + return "ann-" + hex.EncodeToString(b) +} diff --git a/internal/engine/announcements_test.go b/internal/engine/announcements_test.go new file mode 100644 index 00000000..6aeccf94 --- /dev/null +++ b/internal/engine/announcements_test.go @@ -0,0 +1,63 @@ +package engine + +import ( + "testing" + "time" +) + +func TestAnnouncementFeed_PostAndActive(t *testing.T) { + af := NewAnnouncementFeed() + + a1 := af.Post(AnnouncementInfo, "System maintenance in 1 hour", 1*time.Hour) + a2 := af.Post(AnnouncementWarning, "Rate limit approaching", 0) + + active := af.Active() + if len(active) != 2 { + t.Fatalf("expected 2 active announcements, got %d", len(active)) + } + + unread := af.Unread() + if len(unread) != 2 { + t.Fatalf("expected 2 unread announcements, got %d", len(unread)) + } + + if !af.MarkRead(a1.ID) { + t.Error("expected MarkRead to succeed for a1") + } + + unread = af.Unread() + if len(unread) != 1 || unread[0].ID != a2.ID { + t.Errorf("expected 1 unread announcement (a2), got %v", unread) + } + + af.MarkAllRead() + if len(af.Unread()) != 0 { + t.Error("expected 0 unread announcements after MarkAllRead") + } +} + +func TestAnnouncementFeed_Expiration(t *testing.T) { + af := NewAnnouncementFeed() + + // Post with short TTL + af.Post(AnnouncementInfo, "Temporary notice", 10*time.Millisecond) + af.Post(AnnouncementSystem, "Permanent notice", 0) + + time.Sleep(25 * time.Millisecond) + + active := af.Active() + if len(active) != 1 || active[0].Message != "Permanent notice" { + t.Errorf("expected only permanent notice after expiration, got %v", active) + } +} + +func TestAnnouncementFeed_Clear(t *testing.T) { + af := NewAnnouncementFeed() + af.Post(AnnouncementInfo, "Notice 1", 0) + af.Post(AnnouncementInfo, "Notice 2", 0) + + af.Clear() + if len(af.Active()) != 0 { + t.Error("expected 0 announcements after Clear") + } +} diff --git a/internal/engine/persistence_service.go b/internal/engine/persistence_service.go index 6e04a1f7..c917a343 100644 --- a/internal/engine/persistence_service.go +++ b/internal/engine/persistence_service.go @@ -197,6 +197,23 @@ func (s *PersistenceService) AddUser(content string) { s.AppendUserJournaled(types.EyrieMessage{Role: "user", Content: content}) } +// AddMessage appends a message with the specified role and content. +func (s *PersistenceService) AddMessage(role, content string) { + if s == nil { + return + } + switch strings.ToLower(role) { + case "assistant": + s.AddAssistant(content) + case "user": + s.AddUser(content) + default: + s.mu.Lock() + s.messages = append(s.messages, types.EyrieMessage{Role: role, Content: content}) + s.mu.Unlock() + } +} + // AddUserWithImage appends a user message with an inline image. // The image is stored as a data URL ("data:;base64,") // so the LLM-side eyrie client can decode it from the message body diff --git a/internal/engine/prompt_queue.go b/internal/engine/prompt_queue.go new file mode 100644 index 00000000..e8e233c2 --- /dev/null +++ b/internal/engine/prompt_queue.go @@ -0,0 +1,190 @@ +package engine + +import ( + "crypto/rand" + "encoding/hex" + "sort" + "sync" + "time" +) + +// PromptPriority defines scheduling urgency for enqueued prompt turns. +type PromptPriority int + +const ( + // PriorityNormal is standard user turns (default FIFO). + PriorityNormal PromptPriority = 0 + // PrioritySteering is scheduled prompts or background notifications. + PrioritySteering PromptPriority = 10 + // PriorityInterjection is high-priority immediate interjections (/btw). + PriorityInterjection PromptPriority = 20 +) + +// EnqueuedPrompt represents a prompt turn waiting in the queue. +type EnqueuedPrompt struct { + ID string `json:"id"` + Text string `json:"text"` + Priority PromptPriority `json:"priority"` + Source string `json:"source"` + EnqueuedAt time.Time `json:"enqueued_at"` + Metadata map[string]interface{} `json:"metadata,omitempty"` +} + +// PromptQueue provides a thread-safe priority FIFO queue for multi-source prompts. +type PromptQueue struct { + mu sync.Mutex + items []EnqueuedPrompt + paused bool +} + +// NewPromptQueue initializes an empty PromptQueue. +func NewPromptQueue() *PromptQueue { + return &PromptQueue{ + items: make([]EnqueuedPrompt, 0), + } +} + +// Enqueue adds a prompt to the queue, sorted by priority (descending) and enqueue time (ascending). +func (pq *PromptQueue) Enqueue(p EnqueuedPrompt) string { + pq.mu.Lock() + defer pq.mu.Unlock() + + if p.ID == "" { + p.ID = generatePromptID() + } + if p.EnqueuedAt.IsZero() { + p.EnqueuedAt = time.Now() + } + + pq.items = append(pq.items, p) + pq.sortLocked() + return p.ID +} + +// EnqueueText is a helper to enqueue text with a given priority and source. +func (pq *PromptQueue) EnqueueText(text string, priority PromptPriority, source string) string { + return pq.Enqueue(EnqueuedPrompt{ + Text: text, + Priority: priority, + Source: source, + EnqueuedAt: time.Now(), + }) +} + +// Dequeue extracts the highest priority, oldest prompt from the queue. +// Returns false if the queue is empty or currently paused. +func (pq *PromptQueue) Dequeue() (EnqueuedPrompt, bool) { + pq.mu.Lock() + defer pq.mu.Unlock() + + if pq.paused || len(pq.items) == 0 { + return EnqueuedPrompt{}, false + } + + item := pq.items[0] + pq.items = pq.items[1:] + return item, true +} + +// Peek views the next item without removing it. +func (pq *PromptQueue) Peek() (EnqueuedPrompt, bool) { + pq.mu.Lock() + defer pq.mu.Unlock() + + if len(pq.items) == 0 { + return EnqueuedPrompt{}, false + } + return pq.items[0], true +} + +// Len returns the current number of enqueued items. +func (pq *PromptQueue) Len() int { + pq.mu.Lock() + defer pq.mu.Unlock() + return len(pq.items) +} + +// IsEmpty returns true if there are no items in the queue. +func (pq *PromptQueue) IsEmpty() bool { + return pq.Len() == 0 +} + +// Clear drops all pending prompts. +func (pq *PromptQueue) Clear() { + pq.mu.Lock() + defer pq.mu.Unlock() + pq.items = pq.items[:0] +} + +// Drain removes and returns all currently queued items in priority order. +func (pq *PromptQueue) Drain() []EnqueuedPrompt { + pq.mu.Lock() + defer pq.mu.Unlock() + + res := make([]EnqueuedPrompt, len(pq.items)) + copy(res, pq.items) + pq.items = pq.items[:0] + return res +} + +// Pause prevents Dequeue from returning items until Resume is called. +func (pq *PromptQueue) Pause() { + pq.mu.Lock() + defer pq.mu.Unlock() + pq.paused = true +} + +// Resume allows Dequeue to continue processing items. +func (pq *PromptQueue) Resume() { + pq.mu.Lock() + defer pq.mu.Unlock() + pq.paused = false +} + +// IsPaused reports whether the queue is currently paused. +func (pq *PromptQueue) IsPaused() bool { + pq.mu.Lock() + defer pq.mu.Unlock() + return pq.paused +} + +// Remove removes a specific prompt by its ID. +func (pq *PromptQueue) Remove(id string) bool { + pq.mu.Lock() + defer pq.mu.Unlock() + + for i, item := range pq.items { + if item.ID == id { + pq.items = append(pq.items[:i], pq.items[i+1:]...) + return true + } + } + return false +} + +// List returns a copy of all queued items in current execution order. +func (pq *PromptQueue) List() []EnqueuedPrompt { + pq.mu.Lock() + defer pq.mu.Unlock() + + res := make([]EnqueuedPrompt, len(pq.items)) + copy(res, pq.items) + return res +} + +func (pq *PromptQueue) sortLocked() { + sort.SliceStable(pq.items, func(i, j int) bool { + // Higher priority first + if pq.items[i].Priority != pq.items[j].Priority { + return pq.items[i].Priority > pq.items[j].Priority + } + // Earlier enqueue time first (FIFO within same priority) + return pq.items[i].EnqueuedAt.Before(pq.items[j].EnqueuedAt) + }) +} + +func generatePromptID() string { + b := make([]byte, 8) + _, _ = rand.Read(b) + return "pq-" + hex.EncodeToString(b) +} diff --git a/internal/engine/prompt_queue_test.go b/internal/engine/prompt_queue_test.go new file mode 100644 index 00000000..76f5fbda --- /dev/null +++ b/internal/engine/prompt_queue_test.go @@ -0,0 +1,154 @@ +package engine + +import ( + "sync" + "testing" + "time" +) + +func TestPromptQueue_PriorityAndFIFO(t *testing.T) { + pq := NewPromptQueue() + + // Enqueue in mixed order + id1 := pq.EnqueueText("normal 1", PriorityNormal, "user") + time.Sleep(2 * time.Millisecond) + id2 := pq.EnqueueText("normal 2", PriorityNormal, "user") + time.Sleep(2 * time.Millisecond) + id3 := pq.EnqueueText("steering 1", PrioritySteering, "schedule") + time.Sleep(2 * time.Millisecond) + id4 := pq.EnqueueText("interjection 1", PriorityInterjection, "btw") + + if pq.Len() != 4 { + t.Fatalf("expected length 4, got %d", pq.Len()) + } + + // Dequeue 1: should be interjection 1 + p1, ok := pq.Dequeue() + if !ok || p1.ID != id4 || p1.Text != "interjection 1" { + t.Errorf("first dequeue = %v, want interjection 1", p1) + } + + // Dequeue 2: should be steering 1 + p2, ok := pq.Dequeue() + if !ok || p2.ID != id3 || p2.Text != "steering 1" { + t.Errorf("second dequeue = %v, want steering 1", p2) + } + + // Dequeue 3: should be normal 1 (FIFO before normal 2) + p3, ok := pq.Dequeue() + if !ok || p3.ID != id1 || p3.Text != "normal 1" { + t.Errorf("third dequeue = %v, want normal 1", p3) + } + + // Dequeue 4: should be normal 2 + p4, ok := pq.Dequeue() + if !ok || p4.ID != id2 || p4.Text != "normal 2" { + t.Errorf("fourth dequeue = %v, want normal 2", p4) + } + + // Queue should now be empty + if !pq.IsEmpty() { + t.Error("expected empty queue") + } +} + +func TestPromptQueue_PauseAndResume(t *testing.T) { + pq := NewPromptQueue() + pq.EnqueueText("task 1", PriorityNormal, "user") + + pq.Pause() + if !pq.IsPaused() { + t.Error("expected queue to be paused") + } + + // Dequeue while paused should return false + _, ok := pq.Dequeue() + if ok { + t.Error("expected Dequeue to fail while paused") + } + + // Peek should still work while paused + p, ok := pq.Peek() + if !ok || p.Text != "task 1" { + t.Errorf("Peek while paused failed: %v", p) + } + + pq.Resume() + if pq.IsPaused() { + t.Error("expected queue to be resumed") + } + + p, ok = pq.Dequeue() + if !ok || p.Text != "task 1" { + t.Errorf("Dequeue after resume failed: %v", p) + } +} + +func TestPromptQueue_DrainAndClear(t *testing.T) { + pq := NewPromptQueue() + pq.EnqueueText("item 1", PriorityNormal, "user") + pq.EnqueueText("item 2", PrioritySteering, "schedule") + + items := pq.Drain() + if len(items) != 2 { + t.Errorf("Drain returned %d items, want 2", len(items)) + } + if !pq.IsEmpty() { + t.Error("expected queue to be empty after Drain") + } + + pq.EnqueueText("item 3", PriorityNormal, "user") + pq.Clear() + if pq.Len() != 0 { + t.Errorf("expected 0 items after Clear, got %d", pq.Len()) + } +} + +func TestPromptQueue_Remove(t *testing.T) { + pq := NewPromptQueue() + id1 := pq.EnqueueText("item 1", PriorityNormal, "user") + id2 := pq.EnqueueText("item 2", PriorityNormal, "user") + + if !pq.Remove(id1) { + t.Error("expected Remove(id1) to succeed") + } + if pq.Remove("nonexistent-id") { + t.Error("expected Remove on nonexistent ID to fail") + } + + list := pq.List() + if len(list) != 1 || list[0].ID != id2 { + t.Errorf("expected list with item 2, got %v", list) + } +} + +func TestPromptQueue_Concurrent(t *testing.T) { + pq := NewPromptQueue() + var wg sync.WaitGroup + + for i := 0; i < 50; i++ { + wg.Add(1) + go func(n int) { + defer wg.Done() + pq.EnqueueText("prompt", PromptPriority(n%3), "test") + }(i) + } + + wg.Wait() + if pq.Len() != 50 { + t.Errorf("expected length 50 after concurrent enqueues, got %d", pq.Len()) + } + + dequeued := 0 + for { + _, ok := pq.Dequeue() + if !ok { + break + } + dequeued++ + } + + if dequeued != 50 { + t.Errorf("dequeued %d items, want 50", dequeued) + } +} diff --git a/internal/engine/session.go b/internal/engine/session.go index f8678c9d..1b6649de 100644 --- a/internal/engine/session.go +++ b/internal/engine/session.go @@ -142,6 +142,12 @@ type Session struct { // scheduleManager coordinates session-log-backed schedule timers. scheduleManager *schedule.Manager + // promptQueue manages FIFO priority turns and steering turns. + promptQueue *PromptQueue + + // announcements manages in-session broadcasts and notices. + announcements *AnnouncementFeed + // Control plane (product modes) — orthogonal to SpecStage and shellmode. workMode WorkMode isolation IsolationProfile @@ -239,9 +245,31 @@ func NewSessionWithClient(chat ChatClient, provider, model, systemPrompt string, s.life.SetAgentsAccumulator(agentsAccum) s.life.SetLintLoop(NewLintLoop()) s.life.SetTestLoop(NewTestLoop()) + s.promptQueue = NewPromptQueue() + s.announcements = NewAnnouncementFeed() return s } +// PromptQueue returns the session's prompt turn queue. +func (s *Session) PromptQueue() *PromptQueue { + s.mu.RLock() + defer s.mu.RUnlock() + if s.promptQueue == nil { + return nil + } + return s.promptQueue +} + +// Announcements returns the session's announcement feed. +func (s *Session) Announcements() *AnnouncementFeed { + s.mu.RLock() + defer s.mu.RUnlock() + if s.announcements == nil { + return nil + } + return s.announcements +} + // ReattachTransport swaps the LLM client after deployment routing or provider.json changes. // Also reattaches the ChatService so the agent loop's `s.ChatLLM().Stream` // call site picks up the new client (Phase 7 migration). diff --git a/internal/tool/ask_user_test.go b/internal/tool/ask_user_test.go new file mode 100644 index 00000000..acdf8754 --- /dev/null +++ b/internal/tool/ask_user_test.go @@ -0,0 +1,98 @@ +package tool + +import ( + "context" + "encoding/json" + "strings" + "testing" +) + +func TestAskUserQuestionTool_Metadata(t *testing.T) { + tool := AskUserQuestionTool{} + if tool.Name() != "AskUserQuestion" { + t.Errorf("Name() = %q, want AskUserQuestion", tool.Name()) + } + if len(tool.Aliases()) == 0 || tool.Aliases()[0] != "ask_user" { + t.Errorf("Aliases() = %v, want ask_user", tool.Aliases()) + } + params := tool.Parameters() + if params["type"] != "object" { + t.Errorf("Parameters() type = %v, want object", params["type"]) + } +} + +func TestAskUserQuestionTool_Execute_SimpleQuestion(t *testing.T) { + tool := AskUserQuestionTool{} + var askedQuestion string + + tc := &ToolContext{ + AskUserFn: func(q string) (string, error) { + askedQuestion = q + return "User answer", nil + }, + } + ctx := WithToolContext(context.Background(), tc) + + input := json.RawMessage(`{"question":"Which database to use?"}`) + res, err := tool.Execute(ctx, input) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if res != "User answer" { + t.Errorf("res = %q, want 'User answer'", res) + } + if askedQuestion != "Which database to use?" { + t.Errorf("askedQuestion = %q, want 'Which database to use?'", askedQuestion) + } +} + +func TestAskUserQuestionTool_Execute_WithOptions(t *testing.T) { + tool := AskUserQuestionTool{} + var askedQuestion string + + tc := &ToolContext{ + AskUserFn: func(q string) (string, error) { + askedQuestion = q + return "PostgreSQL", nil + }, + } + ctx := WithToolContext(context.Background(), tc) + + input := json.RawMessage(`{ + "question": "Which database?", + "options": ["PostgreSQL", "MySQL", "SQLite"], + "multi_select": false + }`) + res, err := tool.Execute(ctx, input) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if res != "PostgreSQL" { + t.Errorf("res = %q, want 'PostgreSQL'", res) + } + if !strings.Contains(askedQuestion, "Options:") || !strings.Contains(askedQuestion, "PostgreSQL") { + t.Errorf("expected askedQuestion to contain options, got %q", askedQuestion) + } +} + +func TestAskUserQuestionTool_Execute_ValidationErrors(t *testing.T) { + tool := AskUserQuestionTool{} + + // Invalid JSON + _, err := tool.Execute(context.Background(), json.RawMessage(`{invalid`)) + if err == nil { + t.Error("expected error for invalid JSON") + } + + // Empty question + _, err = tool.Execute(context.Background(), json.RawMessage(`{"question":""}`)) + if err == nil { + t.Error("expected error for empty question") + } + + // Unconfigured context + _, err = tool.Execute(context.Background(), json.RawMessage(`{"question":"Test?"}`)) + if err == nil { + t.Error("expected error when ask_user is not configured") + } +} From a99936ee7ce15ad9f79a14914354900c99010ef8 Mon Sep 17 00:00:00 2001 From: Lakshman Patel Date: Fri, 21 Aug 2026 04:01:37 +0530 Subject: [PATCH 2/2] fix(terminal): ensure terminal read drains buffered output on exit - terminal: Track readDone state in Terminal when background reader terminates. - terminal: Update Read to drain buffered bytes properly. - tests: Update TestTerminal_LifecycleAndRead with bounded polling. --- internal/terminal/store.go | 12 ++++++++---- internal/terminal/terminal_test.go | 16 ++++++++++++---- 2 files changed, 20 insertions(+), 8 deletions(-) diff --git a/internal/terminal/store.go b/internal/terminal/store.go index 2e00047c..1b3bd3d2 100644 --- a/internal/terminal/store.go +++ b/internal/terminal/store.go @@ -54,6 +54,7 @@ type Terminal struct { buf bytes.Buffer closed bool alive bool + readDone bool exitCode int } @@ -75,8 +76,7 @@ func (t *Terminal) Send(input string, enter bool) error { return err } -// Read reads up to maxBytes from the buffered terminal output. -// If timeout > 0, it blocks until new output is available or the timeout expires. +// Read reads pending output bytes from the terminal with a bounded wait timeout. func (t *Terminal) Read(maxBytes int, timeout time.Duration) (string, bool, error) { t.mu.Lock() defer t.mu.Unlock() @@ -86,7 +86,7 @@ func (t *Terminal) Read(maxBytes int, timeout time.Duration) (string, bool, erro } // If no data and timeout specified, wait on cond - if t.buf.Len() == 0 && timeout > 0 && t.alive { + if t.buf.Len() == 0 && timeout > 0 && !t.closed && !t.readDone { timer := time.AfterFunc(timeout, func() { t.mu.Lock() t.cond.Broadcast() @@ -94,7 +94,7 @@ func (t *Terminal) Read(maxBytes int, timeout time.Duration) (string, bool, erro }) defer timer.Stop() - for t.buf.Len() == 0 && t.alive && !t.closed { + for t.buf.Len() == 0 && !t.closed && !t.readDone { t.cond.Wait() break } @@ -275,6 +275,10 @@ func (s *Store) Create(ctx context.Context, sessionID, cwd, command string, rows t.mu.Unlock() } if rErr != nil { + t.mu.Lock() + t.readDone = true + t.cond.Broadcast() + t.mu.Unlock() break } } diff --git a/internal/terminal/terminal_test.go b/internal/terminal/terminal_test.go index 028d11a9..9fc3e33f 100644 --- a/internal/terminal/terminal_test.go +++ b/internal/terminal/terminal_test.go @@ -26,10 +26,18 @@ func TestTerminal_LifecycleAndRead(t *testing.T) { t.Errorf("expected branded terminal ID (terminal-), got %s", term.ID) } - // Read output - out, _, err := term.Read(1024, 2*time.Second) - if err != nil { - t.Fatalf("Read failed: %v", err) + // Read output with polling + var out string + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + chunk, _, rErr := term.Read(1024, 200*time.Millisecond) + if rErr != nil { + t.Fatalf("Read failed: %v", rErr) + } + out += chunk + if strings.Contains(out, "hello_hawk") { + break + } } if !strings.Contains(out, "hello_hawk") { t.Errorf("expected output to contain hello_hawk, got %q", out)