返回 DeepSeek-Reasonix
sampling_recovery.go
根目录 / internal / agent / sampling_recovery.go
1 package agent
2
3 import (
4 "context"
5 "errors"
6 "time"
7
8 "reasonix/internal/event"
9 "reasonix/internal/provider"
10 )
11
12 type samplingRecoveryState struct {
13 frozen samplingRequest
14 context contextRecoveryBudget
15 replay reasoningReplayRecoveryBudget
16 output, protocol, missing bool
17 billable *provider.Usage
18 }
19
20 func (a *Agent) samplingDeadline(ctx context.Context) (context.Context, context.CancelFunc, TaskBudget) {
21 limit := a.taskBudgetLimit(ctx)
22 if a.turn.graceRound {
23 limit = TaskBudget{}
24 }
25 if limit.Wall <= 0 {
26 return ctx, func() {}, limit
27 }
28 started := a.task.budget.started
29 if started.IsZero() {
30 started = a.turn.budget.started
31 }
32 if started.IsZero() {
33 started = time.Now()
34 }
35 next, cancel := context.WithDeadline(ctx, started.Add(limit.Wall))
36 return next, cancel, limit
37 }
38
39 func (a *Agent) streamWithSamplingRecovery(parent context.Context, turn int) (terminal streamedTurn) {
40 ctx, cancel, limit := a.samplingDeadline(parent)
41 defer cancel()
42 state := samplingRecoveryState{}
43 defer func() {
44 if limit.Wall > 0 && errors.Is(ctx.Err(), context.DeadlineExceeded) && errors.Is(terminal.err, context.DeadlineExceeded) && parent.Err() == nil {
45 terminal.err = &taskBudgetPause{axis: "time", detail: "recovery reached the task deadline"}
46 }
47 if terminal.err == nil && state.replay.retries > 0 {
48 a.activateReasoningReplayStrongProjection(state.replay)
49 }
50 }()
51 var err error
52 state.frozen, err = a.prepareSamplingRequest(ctx)
53 if err != nil {
54 return streamedTurn{err: err}
55 }
56 if err := a.consumeManualProtocolRecovery(ctx, &state); err != nil {
57 return streamedTurn{err: err}
58 }
59 ctx = provider.WithManagedRecovery(provider.WithRequestAttemptCounter(ctx))
60 for attempt := 1; ; attempt++ {
61 if err := a.samplingRecoveryStop(ctx, limit, state.billable, attempt); err != nil {
62 return streamedTurn{err: err, usage: state.billable}
63 }
64 if state.protocol && !state.replay.persisted {
65 record := a.protocolRecord(state.frozen, "consumed")
66 if state.replay.cutoff > 0 {
67 record.Projected = true
68 record.Prefix, record.Anchor = state.replay.cutoff, state.replay.anchor
69 }
70 if err := a.saveProtocolRecord(record); err != nil {
71 return streamedTurn{err: err, usage: state.billable}
72 }
73 state.replay.persisted = true
74 }
75 id := newStreamAttemptID(attempt)
76 a.emitStreamAttempt(id, event.StreamAttemptBegin, attempt, "", nil)
77 sink, attemptSink := a.samplingAttemptSinks()
78 a.freezeVisibleReads(state.frozen.req.Messages)
79 result := a.runSamplingAttempt(ctx, turn, attemptSink, &state.frozen, id)
80 state.billable, _ = a.recordSamplingAttempt(state.billable, result)
81 if ctx.Err() != nil {
82 // A user cancellation settles the visible prefix as local display
83 // history. Dropping the attempt here loses the only complete prefix.
84 sink.Flush()
85 result.err, result.interrupted, result.usage = ctx.Err(), true, state.billable
86 return result
87 }
88 if result.err == nil {
89 retry, done := a.handleSamplingCandidate(&state, result, sink, attempt, id)
90 if retry {
91 continue
92 }
93 return done
94 }
95 if attempt < maxSamplingAttempts && a.trySamplingRepair(ctx, &state, result, sink, attempt, id) {
96 continue
97 }
98 sink.Flush()
99 if state.context.failure != nil {
100 result.err = state.context.failure
101 }
102 if !state.protocol {
103 if err := a.offerProtocolRecovery(state.frozen, result.err); err != nil {
104 result.err = err
105 }
106 }
107 if provider.AsContextLimitError(result.err) != nil {
108 a.setLastRecovery(contextRecoveryFailed)
109 }
110 result.usage = finalizeSamplingUsage(state.billable, result.usage)
111 if ctx.Err() != nil {
112 result.err = ctx.Err()
113 result.interrupted = true
114 }
115 return result
116 }
117 }
118
119 func (a *Agent) samplingRecoveryStop(ctx context.Context, limit TaskBudget, usage *provider.Usage, attempt int) error {
120 if ctx.Err() != nil {
121 return ctx.Err()
122 }
123 if attempt <= 1 {
124 return nil
125 }
126 shadow := a.task.budget
127 if usage != nil {
128 shadow.observe(usage, a.svc.pricing)
129 }
130 if axis, detail := shadow.exceeded(limit); axis != "" {
131 return &taskBudgetPause{axis: axis, detail: detail}
132 }
133 return nil
134 }
135
136 func (a *Agent) handleSamplingCandidate(s *samplingRecoveryState, result streamedTurn, sink *deferredStreamSink, attempt int, id string) (bool, streamedTurn) {
137 issue := a.reasoningReplayIssue(result)
138 if issue == "" {
139 a.observeMissingAssistantReasoning(result.assistantMessage(), result.reasoningComplete)
140 if s.missing {
141 a.recordRecoveredCandidate(result)
142 }
143 sink.Flush()
144 result.settledAttemptID, result.settledAttempt = id, attempt
145 result.usage = finalizeSamplingUsage(s.billable, result.usage)
146 return false, result
147 }
148 _, claimed := a.observeMissingAssistantReasoning(result.assistantMessage(), result.reasoningComplete)
149 if (issue != ReasoningReplayMissing && issue != ReasoningReplayIncomplete) || s.protocol || a.protocolRecoverySpent() || !claimed || attempt >= maxSamplingAttempts {
150 return false, a.finishReasoningReplayOverflow(result, sink, issue, s.billable, id, attempt)
151 }
152 s.protocol, s.missing = true, true
153 event.RecordProtocolRecovery(a.svc.sink, event.ProtocolRecoveryAudit{Kind: event.ProtocolRecoveryMissingReasoningRetryAttempted})
154 if next, ok := a.recoverReasoningReplayHistory(s.frozen, &s.replay); ok {
155 s.frozen = next
156 s.replay.local = true
157 }
158 sink.Discard()
159 a.emitStreamAttempt(id, event.StreamAttemptDiscard, attempt, "reasoning_replay", nil)
160 a.emitProtocolRetry(attempt, false)
161 return true, streamedTurn{}
162 }
163
164 func (a *Agent) recordRecoveredCandidate(result streamedTurn) {
165 kind := event.ProtocolRecoveryMissingReasoningRetryRecovered
166 if len(result.calls) == 0 && len(result.serverSearch) == 0 {
167 kind = event.ProtocolRecoveryMissingReasoningRetryReplaced
168 }
169 event.RecordProtocolRecovery(a.svc.sink, event.ProtocolRecoveryAudit{Kind: kind})
170 }
171
172 func (a *Agent) trySamplingRepair(ctx context.Context, s *samplingRecoveryState, result streamedTurn, sink *deferredStreamSink, attempt int, id string) bool {
173 if limit := provider.AsOutputLimitError(result.err); !s.output && limit != nil && s.frozen.req.MaxTokens > limit.MaxOutputTokens {
174 s.output = true
175 a.learnOutputBudget(limit.MaxOutputTokens)
176 s.frozen.req.MaxTokens = limit.MaxOutputTokens
177 sink.Discard()
178 a.emitStreamAttempt(id, event.StreamAttemptDiscard, attempt, "output_limit", result.err)
179 return true
180 }
181 if next, ok, _ := a.recoverContextLimit(ctx, s.frozen, result.err, &s.context); ok {
182 sink.Discard()
183 a.emitStreamAttempt(id, event.StreamAttemptDiscard, attempt, "context_limit", result.err)
184 s.frozen = next
185 return true
186 }
187 if s.protocol || s.context.failure != nil {
188 return false
189 }
190 next, ok := a.tryRecoverReasoningReplay400(sink, s.frozen, id, attempt, result.err, &s.replay)
191 if ok {
192 s.protocol = true
193 s.frozen = next
194 }
195 return ok
196 }
197
198 func unmeteredHeaderFailure(result streamedTurn, httpRequests int) bool {
199 if httpRequests <= 0 || sawSpeculativeSamplingOutput(result) {
200 return false
201 }
202 failure := provider.ClassifyRecovery(result.err)
203 return failure.Phase == "headers" || failure.Phase == "connect"
204 }
205
206 func unmeteredUsage(usage *provider.Usage, result streamedTurn, httpRequests int) *provider.Usage {
207 if usage == nil && unmeteredHeaderFailure(result, httpRequests) {
208 return &provider.Usage{Unknown: true, RequestCount: httpRequests}
209 }
210 return usage
211 }
212
212 lines GO