mirror of
https://github.com/coder/coder.git
synced 2026-09-22 21:22:17 +08:00
Implements the chatd stabilization RFC. Combines: - https://github.com/coder/coder/pull/25908 - https://github.com/coder/coder/pull/25923 - https://github.com/coder/coder/pull/26109 - https://github.com/coder/coder/pull/26110 - https://github.com/coder/coder/pull/26111 - https://github.com/coder/coder/pull/26112
249 lines
5.9 KiB
Go
249 lines
5.9 KiB
Go
package chatd
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"net/http"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/google/uuid"
|
|
|
|
"cdr.dev/slog/v3"
|
|
"github.com/coder/quartz"
|
|
)
|
|
|
|
const (
|
|
streamRelayRetryInitialBackoff = 100 * time.Millisecond
|
|
streamRelayRetryMaxBackoff = 5 * time.Second
|
|
)
|
|
|
|
type streamRelayForwarder struct {
|
|
chatID uuid.UUID
|
|
requestHeader http.Header
|
|
dialer StreamPartsDialer
|
|
clock quartz.Clock
|
|
logger slog.Logger
|
|
|
|
parts chan StreamPart
|
|
|
|
ctx context.Context
|
|
cancel context.CancelFunc
|
|
done chan struct{}
|
|
|
|
configure chan streamRelayTarget
|
|
closeOnce sync.Once
|
|
}
|
|
|
|
func newStreamRelayForwarder(
|
|
chatID uuid.UUID,
|
|
requestHeader http.Header,
|
|
dialer StreamPartsDialer,
|
|
clock quartz.Clock,
|
|
logger slog.Logger,
|
|
) *streamRelayForwarder {
|
|
if clock == nil {
|
|
clock = quartz.NewReal()
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
f := &streamRelayForwarder{
|
|
chatID: chatID,
|
|
requestHeader: cloneHeader(requestHeader),
|
|
dialer: dialer,
|
|
clock: clock,
|
|
logger: logger,
|
|
parts: make(chan StreamPart, 128),
|
|
ctx: ctx,
|
|
cancel: cancel,
|
|
done: make(chan struct{}),
|
|
configure: make(chan streamRelayTarget, 1),
|
|
}
|
|
go f.loop()
|
|
return f
|
|
}
|
|
|
|
func (f *streamRelayForwarder) Parts() <-chan StreamPart {
|
|
return f.parts
|
|
}
|
|
|
|
func (f *streamRelayForwarder) Configure(ctx context.Context, target streamRelayTarget) {
|
|
if f == nil {
|
|
return
|
|
}
|
|
// Drop any pending target so the buffered channel always holds the most
|
|
// recent configuration.
|
|
select {
|
|
case <-f.configure:
|
|
default:
|
|
}
|
|
select {
|
|
case f.configure <- target:
|
|
case <-f.ctx.Done():
|
|
case <-ctx.Done():
|
|
}
|
|
}
|
|
|
|
func (f *streamRelayForwarder) Close() {
|
|
if f == nil {
|
|
return
|
|
}
|
|
f.closeOnce.Do(func() {
|
|
f.cancel()
|
|
<-f.done
|
|
})
|
|
}
|
|
|
|
func (f *streamRelayForwarder) loop() {
|
|
defer close(f.done)
|
|
defer close(f.parts)
|
|
var (
|
|
target streamRelayTarget
|
|
connected streamRelayTarget
|
|
session StreamPartsSession
|
|
sessionParts <-chan StreamPart
|
|
retryTimer *quartz.Timer
|
|
retryC <-chan time.Time
|
|
retryBackoff = streamRelayRetryInitialBackoff
|
|
)
|
|
stopRetry := func() {
|
|
if retryTimer != nil {
|
|
retryTimer.Stop()
|
|
retryTimer = nil
|
|
retryC = nil
|
|
}
|
|
}
|
|
defer stopRetry()
|
|
closeSession := func() {
|
|
if session != nil {
|
|
_ = session.Close()
|
|
}
|
|
session = nil
|
|
sessionParts = nil
|
|
connected = streamRelayTarget{}
|
|
}
|
|
defer closeSession()
|
|
scheduleRetry := func() {
|
|
if !target.needsRelay() || f.dialer == nil || retryTimer != nil {
|
|
return
|
|
}
|
|
retryTimer = f.clock.NewTimer(retryBackoff, "chatd", "stream-relay-retry")
|
|
retryC = retryTimer.C
|
|
if retryBackoff < streamRelayRetryMaxBackoff {
|
|
retryBackoff *= 2
|
|
if retryBackoff > streamRelayRetryMaxBackoff {
|
|
retryBackoff = streamRelayRetryMaxBackoff
|
|
}
|
|
}
|
|
}
|
|
connect := func(ctx context.Context) {
|
|
stopRetry()
|
|
if !target.needsRelay() {
|
|
closeSession()
|
|
return
|
|
}
|
|
if f.dialer == nil {
|
|
return
|
|
}
|
|
if session != nil && connected.workerID.Valid && sameNullUUID(connected.workerID, target.workerID) {
|
|
if err := session.SelectEpisode(ctx, target.historyVersion, target.generationAttempt); err != nil {
|
|
f.logger.Warn(ctx, "failed to select stream parts episode",
|
|
slog.F("chat_id", f.chatID),
|
|
slog.F("history_version", target.historyVersion),
|
|
slog.F("generation_attempt", target.generationAttempt),
|
|
slog.Error(err),
|
|
)
|
|
closeSession()
|
|
scheduleRetry()
|
|
return
|
|
}
|
|
connected = target
|
|
retryBackoff = streamRelayRetryInitialBackoff
|
|
return
|
|
}
|
|
closeSession()
|
|
newSession, err := f.dialer(ctx, StreamPartsDialInput{
|
|
ChatID: f.chatID,
|
|
WorkerID: target.workerID.UUID,
|
|
RequestHeader: cloneHeader(f.requestHeader),
|
|
})
|
|
if err != nil {
|
|
f.logger.Warn(ctx, "failed to dial stream parts relay",
|
|
slog.F("chat_id", f.chatID),
|
|
slog.F("worker_id", target.workerID.UUID),
|
|
slog.Error(err),
|
|
)
|
|
// Unrecoverable dial errors (e.g. auth failures) will not
|
|
// succeed on retry with the same inputs, so wait for the next
|
|
// configuration instead of scheduling a retry.
|
|
if !streamPartsDialUnrecoverable(err) {
|
|
scheduleRetry()
|
|
}
|
|
return
|
|
}
|
|
session = newSession
|
|
sessionParts = newSession.Parts()
|
|
connected = streamRelayTarget{workerID: target.workerID}
|
|
if err := session.SelectEpisode(ctx, target.historyVersion, target.generationAttempt); err != nil {
|
|
f.logger.Warn(ctx, "failed to select stream parts episode",
|
|
slog.F("chat_id", f.chatID),
|
|
slog.F("history_version", target.historyVersion),
|
|
slog.F("generation_attempt", target.generationAttempt),
|
|
slog.Error(err),
|
|
)
|
|
closeSession()
|
|
scheduleRetry()
|
|
return
|
|
}
|
|
connected = target
|
|
retryBackoff = streamRelayRetryInitialBackoff
|
|
}
|
|
|
|
for {
|
|
select {
|
|
case <-f.ctx.Done():
|
|
return
|
|
case nextTarget := <-f.configure:
|
|
target = nextTarget
|
|
connect(f.ctx)
|
|
case <-retryC:
|
|
retryTimer = nil
|
|
retryC = nil
|
|
connect(f.ctx)
|
|
case part, ok := <-sessionParts:
|
|
if !ok {
|
|
closeSession()
|
|
scheduleRetry()
|
|
continue
|
|
}
|
|
if !connected.sameEpisode(target) ||
|
|
part.HistoryVersion != target.historyVersion ||
|
|
part.GenerationAttempt != target.generationAttempt {
|
|
continue
|
|
}
|
|
select {
|
|
case f.parts <- part:
|
|
case <-f.ctx.Done():
|
|
return
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
func (t streamRelayTarget) needsRelay() bool {
|
|
return t.workerID.Valid && t.generationAttempt > 0
|
|
}
|
|
|
|
// streamPartsDialUnrecoverable reports whether a dial error signals that
|
|
// retrying with the same inputs is futile, such as an auth failure. Dialers
|
|
// opt in by returning errors that implement IsUnrecoverable.
|
|
func streamPartsDialUnrecoverable(err error) bool {
|
|
var unrecoverable interface{ IsUnrecoverable() bool }
|
|
return errors.As(err, &unrecoverable) && unrecoverable.IsUnrecoverable()
|
|
}
|
|
|
|
func (t streamRelayTarget) sameEpisode(other streamRelayTarget) bool {
|
|
return sameNullUUID(t.workerID, other.workerID) &&
|
|
t.historyVersion == other.historyVersion &&
|
|
t.generationAttempt == other.generationAttempt
|
|
}
|