| 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 |