Skip to content
Open
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
34 changes: 33 additions & 1 deletion agent/codex/appserver_session.go
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,10 @@ type turnStartResponse struct {
} `json:"turn"`
}

type turnSteerResponse struct {
TurnID string `json:"turnId"`
}

type turnNotification struct {
ThreadID string `json:"threadId"`
Turn struct {
Expand Down Expand Up @@ -457,7 +461,8 @@ func (s *appServerSession) Send(prompt string, messageID string, images []core.I
}

s.stateMu.Lock()
if !s.preambleSent {
activeTurn := s.currentTurn
if activeTurn == "" && !s.preambleSent {
prompt = prependCodexPromptPreamble(prompt, s.promptPreamble)
s.preambleSent = true
}
Expand All @@ -481,6 +486,13 @@ func (s *appServerSession) Send(prompt string, messageID string, images []core.I
})
}

if activeTurn != "" {
return s.steerTurn(threadID, activeTurn, input)
}
return s.startTurn(threadID, input)
}

func (s *appServerSession) startTurn(threadID string, input []map[string]any) error {
params := map[string]any{
"threadId": threadID,
"input": input,
Expand Down Expand Up @@ -511,6 +523,26 @@ func (s *appServerSession) Send(prompt string, messageID string, images []core.I
return nil
}

func (s *appServerSession) steerTurn(threadID, expectedTurnID string, input []map[string]any) error {
params := map[string]any{
"threadId": threadID,
"expectedTurnId": expectedTurnID,
"input": input,
}

var resp turnSteerResponse
if err := s.request("turn/steer", params, &resp); err != nil {
return fmt.Errorf("codex app-server turn/steer: %w", err)
}
if resp.TurnID == "" {
return fmt.Errorf("codex app-server turn/steer returned empty turn id")
}
if resp.TurnID != expectedTurnID {
return fmt.Errorf("codex app-server turn/steer returned turn id %q, want %q", resp.TurnID, expectedTurnID)
}
return nil
}

func (s *appServerSession) stageImages(prompt string, images []core.ImageAttachment) (string, []string, error) {
if len(images) == 0 {
return prompt, nil, nil
Expand Down
166 changes: 166 additions & 0 deletions agent/codex/appserver_session_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -177,6 +177,98 @@ func TestAppServerSession_RequestTimeoutIncludesBlockedStdinWrite(t *testing.T)
}
}

func TestAppServerSession_SendSteersActiveTurn(t *testing.T) {
s, stdin := newSendTestSession(t, "turn-1", "buffered answer")
request, err := sendAndRespond(t, s, stdin, "focus on the requested fix", map[string]any{"turnId": "turn-1"})
if err != nil {
t.Fatalf("Send() error = %v", err)
}
if request.Method != "turn/steer" {
t.Fatalf("method = %q, want turn/steer", request.Method)
}
if got := request.Params["threadId"]; got != "thread-1" {
t.Fatalf("threadId = %v, want thread-1", got)
}
if got := request.Params["expectedTurnId"]; got != "turn-1" {
t.Fatalf("expectedTurnId = %v, want turn-1", got)
}
if got := request.Params["input"]; got == nil {
t.Fatal("input = nil, want structured steering input")
}

currentTurn, pendingMsgs := sendTestState(s)
if currentTurn != "turn-1" {
t.Fatalf("current turn = %q, want turn-1", currentTurn)
}
if len(pendingMsgs) != 1 || pendingMsgs[0] != "buffered answer" {
t.Fatalf("pending messages = %v, want buffered answer", pendingMsgs)
}

completedItem, err := json.Marshal(map[string]any{
"threadId": "thread-1",
"turnId": "turn-1",
"item": map[string]any{"type": "agentMessage", "text": "complete final answer"},
})
if err != nil {
t.Fatalf("marshal completed item: %v", err)
}
s.handleNotification("item/completed", completedItem)
completedTurn, err := json.Marshal(map[string]any{
"threadId": "thread-1",
"turn": map[string]any{"id": "turn-1", "status": "completed"},
})
if err != nil {
t.Fatalf("marshal completed turn: %v", err)
}
s.handleNotification("turn/completed", completedTurn)

firstText := waitForEvent(t, s.events)
finalText := waitForEvent(t, s.events)
resultEvent := waitForEvent(t, s.events)
if firstText.Type != core.EventText || firstText.Content != "buffered answer" {
t.Fatalf("first text event = %#v, want buffered answer", firstText)
}
if finalText.Type != core.EventText || finalText.Content != "complete final answer" {
t.Fatalf("final text event = %#v, want complete final answer", finalText)
}
if resultEvent.Type != core.EventResult || !resultEvent.Done {
t.Fatalf("result event = %#v, want one completed result", resultEvent)
}
}

func TestAppServerSession_SendStartsIdleTurn(t *testing.T) {
s, stdin := newSendTestSession(t, "", "stale message")
request, err := sendAndRespond(t, s, stdin, "start new work", map[string]any{"turn": map[string]any{"id": "turn-2"}})
if err != nil {
t.Fatalf("Send() error = %v", err)
}
if request.Method != "turn/start" {
t.Fatalf("method = %q, want turn/start", request.Method)
}
currentTurn, pendingMsgs := sendTestState(s)
if currentTurn != "turn-2" {
t.Fatalf("current turn = %q, want turn-2", currentTurn)
}
if len(pendingMsgs) != 0 {
t.Fatalf("pending messages = %v, want none", pendingMsgs)
}
}

func TestAppServerSession_SendRejectsMismatchedSteeringTurn(t *testing.T) {
s, stdin := newSendTestSession(t, "turn-1", "buffered answer")
_, err := sendAndRespond(t, s, stdin, "steer current work", map[string]any{"turnId": "turn-other"})
if err == nil || !strings.Contains(err.Error(), "turn/steer") || !strings.Contains(err.Error(), "turn-other") {
t.Fatalf("Send() error = %v, want mismatched turn/steer error", err)
}
currentTurn, pendingMsgs := sendTestState(s)
if currentTurn != "turn-1" {
t.Fatalf("current turn = %q, want turn-1", currentTurn)
}
if len(pendingMsgs) != 1 || pendingMsgs[0] != "buffered answer" {
t.Fatalf("pending messages = %v, want buffered answer", pendingMsgs)
}
}

func TestMapAppServerRateLimits_PrefersMultiBucketView(t *testing.T) {
report := mapAppServerRateLimits(appServerRateLimitsResponse{
RateLimits: appServerRateLimitSnapshot{
Expand Down Expand Up @@ -432,6 +524,80 @@ func serverRequestProbe(t *testing.T, idJSON, method string, params any) map[str
}
}

type clientRequestProbe struct {
ID any `json:"id"`
Method string `json:"method"`
Params map[string]any `json:"params"`
}

func newSendTestSession(t *testing.T, currentTurn string, pendingMsgs ...string) (*appServerSession, *lockedWriteCloser) {
t.Helper()
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
stdin := &lockedWriteCloser{}
s := &appServerSession{
ctx: ctx,
cancel: cancel,
events: make(chan core.Event, 8),
stdin: stdin,
pending: make(map[int64]chan rpcResponseEnvelope),
preambleSent: true,
currentTurn: currentTurn,
pendingMsgs: append([]string(nil), pendingMsgs...),
}
s.alive.Store(true)
s.threadID.Store("thread-1")
return s, stdin
}

func sendAndRespond(t *testing.T, s *appServerSession, stdin *lockedWriteCloser, prompt string, result any) (clientRequestProbe, error) {
t.Helper()
done := make(chan error, 1)
go func() {
done <- s.Send(prompt, "", nil, nil)
}()
request := waitForClientRequest(t, stdin)
resultJSON, err := json.Marshal(result)
if err != nil {
t.Fatalf("marshal response: %v", err)
}
s.handleResponse(rpcResponseEnvelope{ID: request.ID, Result: resultJSON})
select {
case err := <-done:
return request, err
case <-time.After(time.Second):
t.Fatal("timed out waiting for Send")
return clientRequestProbe{}, nil
}
}

func sendTestState(s *appServerSession) (string, []string) {
s.stateMu.Lock()
defer s.stateMu.Unlock()
return s.currentTurn, append([]string(nil), s.pendingMsgs...)
}

func waitForEvent(t *testing.T, events <-chan core.Event) core.Event {
t.Helper()
select {
case event := <-events:
return event
case <-time.After(time.Second):
t.Fatal("timed out waiting for app-server event")
return core.Event{}
}
}

func waitForClientRequest(t *testing.T, w *lockedWriteCloser) clientRequestProbe {
t.Helper()
line := waitForWrittenJSONLine(t, w)
var request clientRequestProbe
if err := json.Unmarshal([]byte(line), &request); err != nil {
t.Fatalf("decode JSON-RPC request %q: %v", line, err)
}
return request
}

func waitForWrittenJSONLine(t *testing.T, w *lockedWriteCloser) string {
t.Helper()
deadline := time.After(time.Second)
Expand Down
20 changes: 4 additions & 16 deletions core/cuj_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1138,16 +1138,10 @@ func TestCUJ_A3_ImageReachesAgent(t *testing.T) {
e.ReceiveMessage(plat, msg)

deadline := time.After(2 * time.Second)
for {
agent.mu.Lock()
n := len(agent.sessions)
agent.mu.Unlock()
if n > 0 {
break
}
for len(plat.getSent()) == 0 {
select {
case <-deadline:
t.Fatal("agent never received the message with image")
t.Fatal("agent never completed the message with image")
default:
time.Sleep(10 * time.Millisecond)
}
Expand Down Expand Up @@ -1202,16 +1196,10 @@ func TestCUJ_A5_FileReachesAgent(t *testing.T) {
e.ReceiveMessage(plat, msg)

deadline := time.After(2 * time.Second)
for {
agent.mu.Lock()
n := len(agent.sessions)
agent.mu.Unlock()
if n > 0 {
return
}
for len(plat.getSent()) == 0 {
select {
case <-deadline:
t.Fatal("agent never received the message with file attachment")
t.Fatal("agent never completed the message with file attachment")
default:
time.Sleep(10 * time.Millisecond)
}
Expand Down
18 changes: 9 additions & 9 deletions core/engine_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -634,19 +634,19 @@ func (a *stubDeleteAgent) DeleteSession(_ context.Context, sessionID string) err
return nil
}

// waitDeleteModePhase polls the delete-mode state for the given session key
// until it reaches the target phase or the timeout expires.
func waitDeleteModePhase(t *testing.T, e *Engine, sessionKey, targetPhase string) {
// waitDeleteModeResult polls until the asynchronous delete operation updates
// both its state and the user-visible card.
func waitDeleteModeResult(t *testing.T, e *Engine, p *stubCardPlatform, sessionKey string) {
t.Helper()
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
dm := e.getDeleteModeState(sessionKey)
if dm != nil && dm.phase == targetPhase {
if dm != nil && dm.phase == "result" && len(p.getRefreshedCards()) > 0 {
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatalf("timed out waiting for delete mode phase %q", targetPhase)
t.Fatal("timed out waiting for delete mode result card")
}

type stubProviderAgent struct {
Expand Down Expand Up @@ -4205,7 +4205,7 @@ func TestDeleteMode_ConfirmAndSubmitDeletesSelectedSessions(t *testing.T) {
}
// Submit is now async; the returned card is a "deleting" indicator.
// Wait for the background goroutine to complete and push the result card.
waitDeleteModePhase(t, e, msg.SessionKey, "result")
waitDeleteModeResult(t, e, p, msg.SessionKey)
if got, want := strings.Join(agent.deleted, ","), "session-1,session-3"; got != want {
t.Fatalf("deleted = %q, want %q", got, want)
}
Expand Down Expand Up @@ -4243,7 +4243,7 @@ func TestDeleteMode_SubmitReportsMissingSelectedSessions(t *testing.T) {
t.Fatal("expected deleting card after submit")
}
// Wait for async deletion to complete.
waitDeleteModePhase(t, e, msg.SessionKey, "result")
waitDeleteModeResult(t, e, p, msg.SessionKey)
refreshed := p.getRefreshedCards()
if len(refreshed) == 0 {
t.Fatal("expected refreshed result card via RefreshCard")
Expand Down Expand Up @@ -4345,7 +4345,7 @@ func TestDeleteMode_SubmitBlocksActiveSession(t *testing.T) {
t.Fatal("expected deleting card")
}
// Wait for async deletion to complete.
waitDeleteModePhase(t, e, msg.SessionKey, "result")
waitDeleteModeResult(t, e, p, msg.SessionKey)
if len(agent.deleted) != 0 {
t.Fatalf("deleted = %v, want none", agent.deleted)
}
Expand Down Expand Up @@ -4421,7 +4421,7 @@ func TestDeleteMode_FormSubmitShowsConfirmThenDeletes(t *testing.T) {
t.Fatal("expected deleting card after submit")
}
// Wait for async deletion to complete.
waitDeleteModePhase(t, e, msg.SessionKey, "result")
waitDeleteModeResult(t, e, p, msg.SessionKey)
if got, want := strings.Join(agent.deleted, ","), "session-1,session-3"; got != want {
t.Fatalf("deleted = %q, want %q", got, want)
}
Expand Down
10 changes: 10 additions & 0 deletions daemon/launchd_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,16 @@ func TestLaunchdStatusUsesUserDomainWhenGUIDomainUnavailable(t *testing.T) {
orig := runLaunchctl
t.Cleanup(func() { runLaunchctl = orig })

dir := t.TempDir()
t.Setenv("HOME", dir)
plistPath := launchdPlistPath()
if err := os.MkdirAll(filepath.Dir(plistPath), 0755); err != nil {
t.Fatalf("MkdirAll() error = %v", err)
}
if err := os.WriteFile(plistPath, []byte("plist"), 0644); err != nil {
t.Fatalf("WriteFile() error = %v", err)
}

guiDomain := launchdGUIDomain()
userDomain := launchdUserDomain()
guiTarget := launchdTarget(guiDomain)
Expand Down
Loading
Loading