mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-08-17 03:27:02 +08:00
Merge pull request #221 from qqlzyhello/fix/bidi-append-seq-restart
同 request_id 下 append_seq 重启时不再误判 stale
This commit is contained in:
@@ -2,6 +2,7 @@ package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -37,8 +38,9 @@ func (tracker *appendSequenceTracker) Acquire(ctx context.Context, requestID str
|
||||
if tracker == nil || strings.TrimSpace(requestID) == "" || appendSeq <= 0 {
|
||||
return appendSequenceTicket{}, false, nil
|
||||
}
|
||||
state := tracker.state(strings.TrimSpace(requestID))
|
||||
stale, err := state.acquire(ctx, appendSeq)
|
||||
requestID = strings.TrimSpace(requestID)
|
||||
state := tracker.state(requestID)
|
||||
stale, err := state.acquire(ctx, requestID, appendSeq)
|
||||
if err != nil || stale {
|
||||
return appendSequenceTicket{}, stale, err
|
||||
}
|
||||
@@ -72,7 +74,7 @@ func (tracker *appendSequenceTracker) state(requestID string) *appendSequenceSta
|
||||
return state
|
||||
}
|
||||
|
||||
func (state *appendSequenceState) acquire(ctx context.Context, appendSeq int64) (bool, error) {
|
||||
func (state *appendSequenceState) acquire(ctx context.Context, requestID string, appendSeq int64) (bool, error) {
|
||||
for {
|
||||
state.mu.Lock()
|
||||
now := time.Now().UTC()
|
||||
@@ -83,6 +85,29 @@ func (state *appendSequenceState) acquire(ctx context.Context, appendSeq int64)
|
||||
state.ready = make(chan struct{})
|
||||
}
|
||||
state.updatedAt = now
|
||||
|
||||
// Cursor may reuse the same request_id for a later turn and restart
|
||||
// append_seqno from 1. Accept that as a sequence restart when idle so
|
||||
// tool results are not discarded as stale forever.
|
||||
if appendSeq == 1 && state.next > 1 {
|
||||
if state.processing {
|
||||
ready := state.ready
|
||||
state.mu.Unlock()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return false, ctx.Err()
|
||||
case <-ready:
|
||||
}
|
||||
continue
|
||||
}
|
||||
prevNext := state.next
|
||||
state.next = 1
|
||||
state.processing = true
|
||||
state.mu.Unlock()
|
||||
log.Printf("forwarder reset append sequence request_id=%s previous_next=%d append_seqno=1", requestID, prevNext)
|
||||
return false, nil
|
||||
}
|
||||
|
||||
switch {
|
||||
case appendSeq < state.next:
|
||||
state.mu.Unlock()
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestAppendSequenceTrackerOrderedAcquire(t *testing.T) {
|
||||
tracker := newAppendSequenceTracker()
|
||||
ctx := context.Background()
|
||||
|
||||
ticket1, stale, err := tracker.Acquire(ctx, "req-1", 1)
|
||||
if err != nil || stale {
|
||||
t.Fatalf("seq 1: stale=%v err=%v", stale, err)
|
||||
}
|
||||
ticket1.Release()
|
||||
|
||||
ticket2, stale, err := tracker.Acquire(ctx, "req-1", 2)
|
||||
if err != nil || stale {
|
||||
t.Fatalf("seq 2: stale=%v err=%v", stale, err)
|
||||
}
|
||||
ticket2.Release()
|
||||
}
|
||||
|
||||
func TestAppendSequenceTrackerRestartFromOne(t *testing.T) {
|
||||
tracker := newAppendSequenceTracker()
|
||||
ctx := context.Background()
|
||||
|
||||
for seq := int64(1); seq <= 5; seq++ {
|
||||
ticket, stale, err := tracker.Acquire(ctx, "req-restart", seq)
|
||||
if err != nil || stale {
|
||||
t.Fatalf("seq %d: stale=%v err=%v", seq, stale, err)
|
||||
}
|
||||
ticket.Release()
|
||||
}
|
||||
|
||||
// Same request_id, Cursor restarts append_seqno from 1 for a new turn.
|
||||
ticket, stale, err := tracker.Acquire(ctx, "req-restart", 1)
|
||||
if err != nil {
|
||||
t.Fatalf("restart seq 1 err=%v", err)
|
||||
}
|
||||
if stale {
|
||||
t.Fatal("expected sequence restart on append_seqno=1 after progress")
|
||||
}
|
||||
ticket.Release()
|
||||
|
||||
ticket2, stale, err := tracker.Acquire(ctx, "req-restart", 2)
|
||||
if err != nil || stale {
|
||||
t.Fatalf("post-restart seq 2: stale=%v err=%v", stale, err)
|
||||
}
|
||||
ticket2.Release()
|
||||
}
|
||||
|
||||
func TestAppendSequenceTrackerMidGapStillStale(t *testing.T) {
|
||||
tracker := newAppendSequenceTracker()
|
||||
ctx := context.Background()
|
||||
|
||||
for seq := int64(1); seq <= 3; seq++ {
|
||||
ticket, stale, err := tracker.Acquire(ctx, "req-gap", seq)
|
||||
if err != nil || stale {
|
||||
t.Fatalf("seq %d: stale=%v err=%v", seq, stale, err)
|
||||
}
|
||||
ticket.Release()
|
||||
}
|
||||
|
||||
_, stale, err := tracker.Acquire(ctx, "req-gap", 2)
|
||||
if err != nil {
|
||||
t.Fatalf("gap retransmit err=%v", err)
|
||||
}
|
||||
if !stale {
|
||||
t.Fatal("expected mid-sequence retransmit below next (but not 1) to stay stale")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAppendSequenceTrackerWaitsForInOrder(t *testing.T) {
|
||||
tracker := newAppendSequenceTracker()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
|
||||
ticket1, stale, err := tracker.Acquire(ctx, "req-wait", 1)
|
||||
if err != nil || stale {
|
||||
t.Fatalf("seq 1: stale=%v err=%v", stale, err)
|
||||
}
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
ticket2, stale, err := tracker.Acquire(ctx, "req-wait", 2)
|
||||
if err != nil {
|
||||
done <- err
|
||||
return
|
||||
}
|
||||
if stale {
|
||||
done <- errors.New("seq 2 unexpectedly stale")
|
||||
return
|
||||
}
|
||||
ticket2.Release()
|
||||
done <- nil
|
||||
}()
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
ticket1.Release()
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
if err != nil {
|
||||
t.Fatalf("waiting seq 2 failed: %v", err)
|
||||
}
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for ordered seq 2")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAppendSequenceTrackerRestartWaitsUntilIdle(t *testing.T) {
|
||||
tracker := newAppendSequenceTracker()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// Advance to next=3.
|
||||
for seq := int64(1); seq <= 2; seq++ {
|
||||
ticket, stale, err := tracker.Acquire(ctx, "req-idle", seq)
|
||||
if err != nil || stale {
|
||||
t.Fatalf("seq %d: stale=%v err=%v", seq, stale, err)
|
||||
}
|
||||
ticket.Release()
|
||||
}
|
||||
|
||||
// Hold seq 3 so a concurrent restart must wait.
|
||||
hold, stale, err := tracker.Acquire(ctx, "req-idle", 3)
|
||||
if err != nil || stale {
|
||||
t.Fatalf("seq 3 hold: stale=%v err=%v", stale, err)
|
||||
}
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
ticket, stale, err := tracker.Acquire(ctx, "req-idle", 1)
|
||||
if err != nil {
|
||||
done <- err
|
||||
return
|
||||
}
|
||||
if stale {
|
||||
done <- errors.New("restart seq 1 unexpectedly stale")
|
||||
return
|
||||
}
|
||||
ticket.Release()
|
||||
done <- nil
|
||||
}()
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
hold.Release()
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
if err != nil {
|
||||
t.Fatalf("restart while busy failed: %v", err)
|
||||
}
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for restart after idle")
|
||||
}
|
||||
|
||||
ticket2, stale, err := tracker.Acquire(ctx, "req-idle", 2)
|
||||
if err != nil || stale {
|
||||
t.Fatalf("post-restart seq 2: stale=%v err=%v", stale, err)
|
||||
}
|
||||
ticket2.Release()
|
||||
}
|
||||
Reference in New Issue
Block a user