Add projects
This commit is contained in:
parent
2d3a9ad623
commit
8b607dd700
1802 changed files with 503346 additions and 2 deletions
818
backend/internal/runtime/executor/codex_websockets_session.go
Normal file
818
backend/internal/runtime/executor/codex_websockets_session.go
Normal file
|
|
@ -0,0 +1,818 @@
|
|||
package executor
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
cliproxyauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth"
|
||||
cliproxyexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
type codexWebsocketSessionStore struct {
|
||||
mu sync.Mutex
|
||||
sessions map[string]*codexWebsocketSession
|
||||
}
|
||||
|
||||
var globalCodexWebsocketSessionStore = &codexWebsocketSessionStore{
|
||||
sessions: make(map[string]*codexWebsocketSession),
|
||||
}
|
||||
|
||||
type websocketConnectionCloser struct {
|
||||
conn *websocket.Conn
|
||||
once sync.Once
|
||||
err error
|
||||
}
|
||||
|
||||
func newWebsocketConnectionCloser(conn *websocket.Conn) *websocketConnectionCloser {
|
||||
if conn == nil {
|
||||
return nil
|
||||
}
|
||||
return &websocketConnectionCloser{conn: conn}
|
||||
}
|
||||
|
||||
func (c *websocketConnectionCloser) Close() error {
|
||||
if c == nil || c.conn == nil {
|
||||
return nil
|
||||
}
|
||||
c.once.Do(func() {
|
||||
c.err = c.conn.Close()
|
||||
})
|
||||
return c.err
|
||||
}
|
||||
|
||||
type codexWebsocketSession struct {
|
||||
sessionID string
|
||||
|
||||
reqMu sync.Mutex
|
||||
|
||||
connMu sync.Mutex
|
||||
conn *websocket.Conn
|
||||
connCloser *websocketConnectionCloser
|
||||
wsURL string
|
||||
authID string
|
||||
multiAgentV2OptimizedConn *websocket.Conn
|
||||
lifecycleBindMu sync.Mutex
|
||||
lifecycle cliproxyexecutor.ExecutionLifecycle
|
||||
lifecycleModel string
|
||||
|
||||
writeMu sync.Mutex
|
||||
|
||||
activeMu sync.Mutex
|
||||
activeConn *websocket.Conn
|
||||
activeCh chan codexWebsocketRead
|
||||
activeDone <-chan struct{}
|
||||
activeCancel context.CancelFunc
|
||||
|
||||
readerConn *websocket.Conn
|
||||
|
||||
upstreamDisconnectOnce sync.Once
|
||||
upstreamDisconnectCh chan error
|
||||
upstreamDisconnectErrMu sync.RWMutex
|
||||
upstreamDisconnectErrConn *websocket.Conn
|
||||
upstreamDisconnectErr error
|
||||
}
|
||||
|
||||
type codexWebsocketRead struct {
|
||||
conn *websocket.Conn
|
||||
msgType int
|
||||
payload []byte
|
||||
err error
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) setActive(conn *websocket.Conn, ch chan codexWebsocketRead) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.activeMu.Lock()
|
||||
if s.activeCancel != nil {
|
||||
s.activeCancel()
|
||||
s.activeCancel = nil
|
||||
s.activeDone = nil
|
||||
}
|
||||
s.activeConn = conn
|
||||
s.activeCh = ch
|
||||
if conn != nil && ch != nil {
|
||||
activeCtx, activeCancel := context.WithCancel(context.Background())
|
||||
s.activeDone = activeCtx.Done()
|
||||
s.activeCancel = activeCancel
|
||||
}
|
||||
s.activeMu.Unlock()
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) activate(conn *websocket.Conn) chan codexWebsocketRead {
|
||||
if s == nil || conn == nil {
|
||||
return nil
|
||||
}
|
||||
ch := make(chan codexWebsocketRead, 4096)
|
||||
s.setActive(conn, ch)
|
||||
return ch
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) activeForConn(conn *websocket.Conn) (chan codexWebsocketRead, <-chan struct{}) {
|
||||
if s == nil || conn == nil {
|
||||
return nil, nil
|
||||
}
|
||||
s.activeMu.Lock()
|
||||
defer s.activeMu.Unlock()
|
||||
if s.activeConn != conn {
|
||||
return nil, nil
|
||||
}
|
||||
return s.activeCh, s.activeDone
|
||||
}
|
||||
|
||||
func clearRetryActiveState(sess *codexWebsocketSession, conn *websocket.Conn, ch chan codexWebsocketRead) bool {
|
||||
if sess == nil {
|
||||
return false
|
||||
}
|
||||
return sess.clearActive(conn, ch)
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) clearActive(conn *websocket.Conn, ch chan codexWebsocketRead) bool {
|
||||
if s == nil {
|
||||
return false
|
||||
}
|
||||
s.activeMu.Lock()
|
||||
defer s.activeMu.Unlock()
|
||||
if s.activeConn != conn || s.activeCh != ch {
|
||||
return false
|
||||
}
|
||||
s.activeConn = nil
|
||||
s.activeCh = nil
|
||||
if s.activeCancel != nil {
|
||||
s.activeCancel()
|
||||
}
|
||||
s.activeCancel = nil
|
||||
s.activeDone = nil
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) writeMessage(conn *websocket.Conn, msgType int, payload []byte) error {
|
||||
if s == nil {
|
||||
return fmt.Errorf("codex websockets executor: session is nil")
|
||||
}
|
||||
if conn == nil {
|
||||
return fmt.Errorf("codex websockets executor: websocket conn is nil")
|
||||
}
|
||||
s.writeMu.Lock()
|
||||
defer s.writeMu.Unlock()
|
||||
return conn.WriteMessage(msgType, payload)
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) setMultiAgentV2Optimized(conn *websocket.Conn, optimized bool) {
|
||||
if s == nil || conn == nil {
|
||||
return
|
||||
}
|
||||
s.connMu.Lock()
|
||||
if s.conn == conn {
|
||||
if optimized {
|
||||
s.multiAgentV2OptimizedConn = conn
|
||||
} else {
|
||||
s.multiAgentV2OptimizedConn = nil
|
||||
}
|
||||
}
|
||||
s.connMu.Unlock()
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) isMultiAgentV2Optimized(conn *websocket.Conn) bool {
|
||||
if s == nil || conn == nil {
|
||||
return false
|
||||
}
|
||||
s.connMu.Lock()
|
||||
defer s.connMu.Unlock()
|
||||
return s.conn == conn && s.multiAgentV2OptimizedConn == conn
|
||||
}
|
||||
|
||||
// sendTerminalWebsocketRead reports whether it invalidated a full channel's connection before waiting.
|
||||
func sendTerminalWebsocketRead(ch chan<- codexWebsocketRead, done <-chan struct{}, event codexWebsocketRead, invalidate func()) bool {
|
||||
select {
|
||||
case ch <- event:
|
||||
return false
|
||||
case <-done:
|
||||
return false
|
||||
default:
|
||||
}
|
||||
|
||||
invalidated := invalidate != nil
|
||||
if invalidated {
|
||||
invalidate()
|
||||
}
|
||||
select {
|
||||
case ch <- event:
|
||||
case <-done:
|
||||
}
|
||||
return invalidated
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) configureConn(conn *websocket.Conn) {
|
||||
if s == nil || conn == nil {
|
||||
return
|
||||
}
|
||||
s.resetUpstreamDisconnectError(conn)
|
||||
conn.SetPingHandler(func(appData string) error {
|
||||
s.writeMu.Lock()
|
||||
defer s.writeMu.Unlock()
|
||||
// Reply pongs from the same write lock to avoid concurrent writes.
|
||||
return conn.WriteControl(websocket.PongMessage, []byte(appData), time.Now().Add(10*time.Second))
|
||||
})
|
||||
defaultCloseHandler := conn.CloseHandler()
|
||||
conn.SetCloseHandler(func(code int, text string) error {
|
||||
s.setUpstreamDisconnectError(conn, &websocket.CloseError{Code: code, Text: text})
|
||||
return defaultCloseHandler(code, text)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) bindExecutionLifecycle(opts cliproxyexecutor.Options, conn *websocket.Conn, closer *websocketConnectionCloser, model string) error {
|
||||
if closer == nil {
|
||||
return fmt.Errorf("codex websockets executor: websocket connection closer is nil")
|
||||
}
|
||||
if s == nil {
|
||||
return cliproxyexecutor.BindExecutionResource(opts, closer)
|
||||
}
|
||||
lifecycle := opts.ExecutionLifecycle
|
||||
if lifecycle == nil || conn == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
s.lifecycleBindMu.Lock()
|
||||
defer s.lifecycleBindMu.Unlock()
|
||||
|
||||
s.connMu.Lock()
|
||||
if s.conn == conn && s.connCloser == nil {
|
||||
s.connCloser = closer
|
||||
}
|
||||
alreadyBound := s.conn == conn && s.connCloser == closer && s.lifecycle == lifecycle
|
||||
s.connMu.Unlock()
|
||||
if alreadyBound {
|
||||
return nil
|
||||
}
|
||||
|
||||
if errBind := lifecycle.Bind(func() error {
|
||||
return s.closeBoundConnection(conn, closer, lifecycle)
|
||||
}); errBind != nil {
|
||||
return errBind
|
||||
}
|
||||
if retained, ok := lifecycle.(interface{ Retain() }); ok {
|
||||
retained.Retain()
|
||||
}
|
||||
|
||||
s.connMu.Lock()
|
||||
if s.conn != conn || s.connCloser != closer {
|
||||
s.connMu.Unlock()
|
||||
return fmt.Errorf("codex websockets executor: websocket connection closed during lifecycle bind")
|
||||
}
|
||||
previous := s.lifecycle
|
||||
s.lifecycle = lifecycle
|
||||
s.lifecycleModel = strings.TrimSpace(model)
|
||||
s.connMu.Unlock()
|
||||
if previous != nil && previous != lifecycle {
|
||||
previous.End("target_replaced")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) closeBoundConnection(conn *websocket.Conn, closer *websocketConnectionCloser, lifecycle cliproxyexecutor.ExecutionLifecycle) error {
|
||||
if s == nil || conn == nil {
|
||||
return nil
|
||||
}
|
||||
s.detachConnection(conn, lifecycle)
|
||||
errClose := closer.Close()
|
||||
go lifecycle.End("connection_closed")
|
||||
return errClose
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) detachConnection(conn *websocket.Conn, lifecycle cliproxyexecutor.ExecutionLifecycle) *websocketConnectionCloser {
|
||||
if s == nil || conn == nil {
|
||||
return nil
|
||||
}
|
||||
s.connMu.Lock()
|
||||
var closer *websocketConnectionCloser
|
||||
matched := s.conn == conn
|
||||
if matched {
|
||||
closer = s.connCloser
|
||||
s.conn = nil
|
||||
s.connCloser = nil
|
||||
s.multiAgentV2OptimizedConn = nil
|
||||
if s.readerConn == conn {
|
||||
s.readerConn = nil
|
||||
}
|
||||
}
|
||||
if (lifecycle == nil && matched) || (lifecycle != nil && s.lifecycle == lifecycle) {
|
||||
s.lifecycle = nil
|
||||
s.lifecycleModel = ""
|
||||
}
|
||||
s.connMu.Unlock()
|
||||
return closer
|
||||
}
|
||||
|
||||
func closeWebsocketAfterBindFailure(sess *codexWebsocketSession, conn *websocket.Conn, closer *websocketConnectionCloser) {
|
||||
if conn == nil || closer == nil {
|
||||
return
|
||||
}
|
||||
if sess != nil {
|
||||
sess.detachConnection(conn, nil)
|
||||
}
|
||||
if errClose := closer.Close(); errClose != nil {
|
||||
log.Errorf("websockets executor: close lifecycle bind failure connection error: %v", errClose)
|
||||
}
|
||||
}
|
||||
|
||||
func websocketSessionTargetChanged(sess *codexWebsocketSession, authID string, wsURL string) bool {
|
||||
if sess == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
sess.connMu.Lock()
|
||||
defer sess.connMu.Unlock()
|
||||
if strings.TrimSpace(sess.authID) == "" && strings.TrimSpace(sess.wsURL) == "" {
|
||||
return false
|
||||
}
|
||||
return strings.TrimSpace(sess.authID) != strings.TrimSpace(authID) || strings.TrimSpace(sess.wsURL) != strings.TrimSpace(wsURL)
|
||||
}
|
||||
|
||||
func existingWebsocketSessionConn(sess *codexWebsocketSession, authID string, wsURL string) (*websocket.Conn, *websocketConnectionCloser) {
|
||||
if sess == nil {
|
||||
return nil, nil
|
||||
}
|
||||
sess.connMu.Lock()
|
||||
conn := sess.conn
|
||||
closer := sess.connCloser
|
||||
matches := conn != nil && closer != nil &&
|
||||
strings.TrimSpace(sess.authID) == strings.TrimSpace(authID) &&
|
||||
strings.TrimSpace(sess.wsURL) == strings.TrimSpace(wsURL)
|
||||
sess.connMu.Unlock()
|
||||
if !matches || sess.upstreamDisconnectError(conn) != nil {
|
||||
return nil, nil
|
||||
}
|
||||
return conn, closer
|
||||
}
|
||||
|
||||
func detachMismatchedWebsocketSessionConn(sess *codexWebsocketSession, authID string, wsURL string) (*websocket.Conn, *websocketConnectionCloser, string, string, cliproxyexecutor.ExecutionLifecycle) {
|
||||
if sess == nil {
|
||||
return nil, nil, "", "", nil
|
||||
}
|
||||
|
||||
sess.connMu.Lock()
|
||||
defer sess.connMu.Unlock()
|
||||
conn := sess.conn
|
||||
if conn == nil || (strings.TrimSpace(sess.authID) == strings.TrimSpace(authID) && strings.TrimSpace(sess.wsURL) == strings.TrimSpace(wsURL)) {
|
||||
return nil, nil, "", "", nil
|
||||
}
|
||||
|
||||
previousAuthID := sess.authID
|
||||
previousWSURL := sess.wsURL
|
||||
lifecycle := sess.lifecycle
|
||||
closer := sess.connCloser
|
||||
sess.lifecycle = nil
|
||||
sess.lifecycleModel = ""
|
||||
sess.conn = nil
|
||||
sess.connCloser = nil
|
||||
sess.multiAgentV2OptimizedConn = nil
|
||||
if sess.readerConn == conn {
|
||||
sess.readerConn = nil
|
||||
}
|
||||
return conn, closer, previousAuthID, previousWSURL, lifecycle
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) resetUpstreamDisconnectError(conn *websocket.Conn) {
|
||||
if s == nil || conn == nil {
|
||||
return
|
||||
}
|
||||
s.upstreamDisconnectErrMu.Lock()
|
||||
s.upstreamDisconnectErrConn = conn
|
||||
s.upstreamDisconnectErr = nil
|
||||
s.upstreamDisconnectErrMu.Unlock()
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) setUpstreamDisconnectError(conn *websocket.Conn, err error) {
|
||||
if s == nil || conn == nil || err == nil {
|
||||
return
|
||||
}
|
||||
s.upstreamDisconnectErrMu.Lock()
|
||||
if s.upstreamDisconnectErrConn == conn && s.upstreamDisconnectErr == nil {
|
||||
s.upstreamDisconnectErr = err
|
||||
}
|
||||
s.upstreamDisconnectErrMu.Unlock()
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) upstreamDisconnectError(conn *websocket.Conn) error {
|
||||
if s == nil || conn == nil {
|
||||
return nil
|
||||
}
|
||||
s.upstreamDisconnectErrMu.RLock()
|
||||
defer s.upstreamDisconnectErrMu.RUnlock()
|
||||
if s.upstreamDisconnectErrConn != conn {
|
||||
return nil
|
||||
}
|
||||
return s.upstreamDisconnectErr
|
||||
}
|
||||
|
||||
func (s *codexWebsocketSession) notifyUpstreamDisconnect(err error) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.upstreamDisconnectOnce.Do(func() {
|
||||
if s.upstreamDisconnectCh == nil {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case s.upstreamDisconnectCh <- err:
|
||||
default:
|
||||
}
|
||||
close(s.upstreamDisconnectCh)
|
||||
})
|
||||
}
|
||||
|
||||
func executionSessionIDFromOptions(opts cliproxyexecutor.Options) string {
|
||||
if len(opts.Metadata) == 0 {
|
||||
return ""
|
||||
}
|
||||
raw, ok := opts.Metadata[cliproxyexecutor.ExecutionSessionMetadataKey]
|
||||
if !ok || raw == nil {
|
||||
return ""
|
||||
}
|
||||
switch v := raw.(type) {
|
||||
case string:
|
||||
return strings.TrimSpace(v)
|
||||
case []byte:
|
||||
return strings.TrimSpace(string(v))
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func (e *CodexWebsocketsExecutor) getOrCreateSession(sessionID string) *codexWebsocketSession {
|
||||
sessionID = strings.TrimSpace(sessionID)
|
||||
if sessionID == "" {
|
||||
return nil
|
||||
}
|
||||
if e == nil {
|
||||
return nil
|
||||
}
|
||||
store := e.store
|
||||
if store == nil {
|
||||
store = globalCodexWebsocketSessionStore
|
||||
}
|
||||
store.mu.Lock()
|
||||
defer store.mu.Unlock()
|
||||
if store.sessions == nil {
|
||||
store.sessions = make(map[string]*codexWebsocketSession)
|
||||
}
|
||||
if sess, ok := store.sessions[sessionID]; ok && sess != nil {
|
||||
return sess
|
||||
}
|
||||
sess := &codexWebsocketSession{
|
||||
sessionID: sessionID,
|
||||
upstreamDisconnectCh: make(chan error, 1),
|
||||
}
|
||||
store.sessions[sessionID] = sess
|
||||
return sess
|
||||
}
|
||||
|
||||
func (e *CodexWebsocketsExecutor) UpstreamDisconnectChan(sessionID string) <-chan error {
|
||||
sess := e.getOrCreateSession(sessionID)
|
||||
if sess == nil {
|
||||
return nil
|
||||
}
|
||||
return sess.upstreamDisconnectCh
|
||||
}
|
||||
|
||||
func (e *CodexWebsocketsExecutor) ensureUpstreamConn(ctx context.Context, auth *cliproxyauth.Auth, sess *codexWebsocketSession, authID string, wsURL string, headers http.Header) (*websocket.Conn, *websocketConnectionCloser, *http.Response, error) {
|
||||
if sess == nil {
|
||||
return e.dialCodexWebsocket(ctx, auth, wsURL, headers)
|
||||
}
|
||||
|
||||
if staleConn, staleCloser, staleAuthID, staleWSURL, staleLifecycle := detachMismatchedWebsocketSessionConn(sess, authID, wsURL); staleConn != nil {
|
||||
logCodexWebsocketDisconnected(sess.sessionID, staleAuthID, staleWSURL, "target_changed", nil)
|
||||
if staleCloser != nil {
|
||||
if errClose := staleCloser.Close(); errClose != nil {
|
||||
log.Errorf("codex websockets executor: close stale websocket error: %v", errClose)
|
||||
}
|
||||
}
|
||||
if staleLifecycle != nil {
|
||||
staleLifecycle.End("target_changed")
|
||||
}
|
||||
}
|
||||
|
||||
sess.connMu.Lock()
|
||||
conn := sess.conn
|
||||
closer := sess.connCloser
|
||||
readerConn := sess.readerConn
|
||||
sess.connMu.Unlock()
|
||||
if conn != nil {
|
||||
if readerConn != conn {
|
||||
sess.connMu.Lock()
|
||||
sess.readerConn = conn
|
||||
sess.connMu.Unlock()
|
||||
sess.configureConn(conn)
|
||||
go e.readUpstreamLoop(sess, conn)
|
||||
}
|
||||
return conn, closer, nil, nil
|
||||
}
|
||||
|
||||
conn, closer, resp, errDial := e.dialCodexWebsocket(ctx, auth, wsURL, headers)
|
||||
if errDial != nil {
|
||||
return nil, closer, resp, errDial
|
||||
}
|
||||
|
||||
sess.connMu.Lock()
|
||||
if sess.conn != nil {
|
||||
previous := sess.conn
|
||||
previousCloser := sess.connCloser
|
||||
sess.connMu.Unlock()
|
||||
if errClose := closer.Close(); errClose != nil {
|
||||
log.Errorf("codex websockets executor: close websocket error: %v", errClose)
|
||||
}
|
||||
return previous, previousCloser, nil, nil
|
||||
}
|
||||
sess.conn = conn
|
||||
sess.connCloser = closer
|
||||
sess.multiAgentV2OptimizedConn = nil
|
||||
sess.wsURL = wsURL
|
||||
sess.authID = authID
|
||||
sess.readerConn = conn
|
||||
sess.connMu.Unlock()
|
||||
|
||||
sess.configureConn(conn)
|
||||
go e.readUpstreamLoop(sess, conn)
|
||||
logCodexWebsocketConnected(sess.sessionID, authID, wsURL)
|
||||
return conn, closer, resp, nil
|
||||
}
|
||||
|
||||
func (e *CodexWebsocketsExecutor) readUpstreamLoop(sess *codexWebsocketSession, conn *websocket.Conn) {
|
||||
if e == nil || sess == nil || conn == nil {
|
||||
return
|
||||
}
|
||||
for {
|
||||
_ = conn.SetReadDeadline(time.Now().Add(codexResponsesWebsocketIdleTimeout))
|
||||
msgType, payload, errRead := conn.ReadMessage()
|
||||
if errRead != nil {
|
||||
invalidate := func() {
|
||||
e.invalidateUpstreamConn(sess, conn, "upstream_disconnected", errRead)
|
||||
}
|
||||
invalidated := false
|
||||
ch, done := sess.activeForConn(conn)
|
||||
if ch != nil {
|
||||
invalidated = sendTerminalWebsocketRead(ch, done, codexWebsocketRead{conn: conn, err: errRead}, invalidate)
|
||||
if sess.clearActive(conn, ch) {
|
||||
close(ch)
|
||||
}
|
||||
}
|
||||
if !invalidated {
|
||||
invalidate()
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
if msgType != websocket.TextMessage {
|
||||
if msgType == websocket.BinaryMessage {
|
||||
errBinary := fmt.Errorf("codex websockets executor: unexpected binary message")
|
||||
invalidate := func() {
|
||||
e.invalidateUpstreamConn(sess, conn, "unexpected_binary", errBinary)
|
||||
}
|
||||
invalidated := false
|
||||
ch, done := sess.activeForConn(conn)
|
||||
if ch != nil {
|
||||
invalidated = sendTerminalWebsocketRead(ch, done, codexWebsocketRead{conn: conn, err: errBinary}, invalidate)
|
||||
if sess.clearActive(conn, ch) {
|
||||
close(ch)
|
||||
}
|
||||
}
|
||||
if !invalidated {
|
||||
invalidate()
|
||||
}
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
ch, done := sess.activeForConn(conn)
|
||||
if ch == nil {
|
||||
continue
|
||||
}
|
||||
select {
|
||||
case ch <- codexWebsocketRead{conn: conn, msgType: msgType, payload: payload}:
|
||||
case <-done:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (e *CodexWebsocketsExecutor) invalidateUpstreamConn(sess *codexWebsocketSession, conn *websocket.Conn, reason string, err error) {
|
||||
e.invalidateUpstreamConnWithNotify(sess, conn, reason, err, true)
|
||||
}
|
||||
|
||||
func (e *CodexWebsocketsExecutor) invalidateUpstreamConnWithoutDisconnectNotify(sess *codexWebsocketSession, conn *websocket.Conn, reason string, err error) {
|
||||
e.invalidateUpstreamConnWithNotify(sess, conn, reason, err, false)
|
||||
}
|
||||
|
||||
func (e *CodexWebsocketsExecutor) invalidateUpstreamConnWithNotify(sess *codexWebsocketSession, conn *websocket.Conn, reason string, err error, notify bool) {
|
||||
if sess == nil || conn == nil {
|
||||
return
|
||||
}
|
||||
|
||||
sess.connMu.Lock()
|
||||
current := sess.conn
|
||||
authID := sess.authID
|
||||
wsURL := sess.wsURL
|
||||
sessionID := sess.sessionID
|
||||
if current == nil || current != conn {
|
||||
sess.connMu.Unlock()
|
||||
return
|
||||
}
|
||||
lifecycle := sess.lifecycle
|
||||
closer := sess.connCloser
|
||||
sess.lifecycle = nil
|
||||
sess.lifecycleModel = ""
|
||||
sess.conn = nil
|
||||
sess.connCloser = nil
|
||||
sess.multiAgentV2OptimizedConn = nil
|
||||
if sess.readerConn == conn {
|
||||
sess.readerConn = nil
|
||||
}
|
||||
sess.connMu.Unlock()
|
||||
|
||||
logCodexWebsocketDisconnected(sessionID, authID, wsURL, reason, err)
|
||||
if notify {
|
||||
sess.notifyUpstreamDisconnect(err)
|
||||
}
|
||||
if closer != nil {
|
||||
if errClose := closer.Close(); errClose != nil {
|
||||
log.Errorf("codex websockets executor: close websocket error: %v", errClose)
|
||||
}
|
||||
}
|
||||
if lifecycle != nil {
|
||||
lifecycle.End(reason)
|
||||
}
|
||||
}
|
||||
|
||||
func (e *CodexWebsocketsExecutor) CloseExecutionSession(sessionID string) {
|
||||
sessionID = strings.TrimSpace(sessionID)
|
||||
if e == nil {
|
||||
return
|
||||
}
|
||||
if sessionID == "" {
|
||||
return
|
||||
}
|
||||
if sessionID == cliproxyauth.CloseAllExecutionSessionsID {
|
||||
e.closeAllExecutionSessions("executor_shutdown")
|
||||
return
|
||||
}
|
||||
|
||||
store := e.store
|
||||
if store == nil {
|
||||
store = globalCodexWebsocketSessionStore
|
||||
}
|
||||
store.mu.Lock()
|
||||
sess := store.sessions[sessionID]
|
||||
delete(store.sessions, sessionID)
|
||||
store.mu.Unlock()
|
||||
|
||||
e.closeExecutionSession(sess, "session_closed")
|
||||
}
|
||||
|
||||
func (e *CodexWebsocketsExecutor) closeAllExecutionSessions(reason string) {
|
||||
if e == nil {
|
||||
return
|
||||
}
|
||||
|
||||
store := e.store
|
||||
if store == nil {
|
||||
store = globalCodexWebsocketSessionStore
|
||||
}
|
||||
store.mu.Lock()
|
||||
sessions := make([]*codexWebsocketSession, 0, len(store.sessions))
|
||||
for sessionID, sess := range store.sessions {
|
||||
delete(store.sessions, sessionID)
|
||||
if sess != nil {
|
||||
sessions = append(sessions, sess)
|
||||
}
|
||||
}
|
||||
store.mu.Unlock()
|
||||
|
||||
for i := range sessions {
|
||||
e.closeExecutionSession(sessions[i], reason)
|
||||
}
|
||||
}
|
||||
|
||||
func (e *CodexWebsocketsExecutor) closeExecutionSession(sess *codexWebsocketSession, reason string) {
|
||||
closeCodexWebsocketSession(sess, reason)
|
||||
}
|
||||
|
||||
func closeCodexWebsocketSession(sess *codexWebsocketSession, reason string) {
|
||||
if sess == nil {
|
||||
return
|
||||
}
|
||||
reason = strings.TrimSpace(reason)
|
||||
if reason == "" {
|
||||
reason = "session_closed"
|
||||
}
|
||||
|
||||
sess.connMu.Lock()
|
||||
conn := sess.conn
|
||||
authID := sess.authID
|
||||
wsURL := sess.wsURL
|
||||
lifecycle := sess.lifecycle
|
||||
closer := sess.connCloser
|
||||
sess.lifecycle = nil
|
||||
sess.lifecycleModel = ""
|
||||
sess.conn = nil
|
||||
sess.connCloser = nil
|
||||
sess.multiAgentV2OptimizedConn = nil
|
||||
if sess.readerConn == conn {
|
||||
sess.readerConn = nil
|
||||
}
|
||||
sessionID := sess.sessionID
|
||||
sess.connMu.Unlock()
|
||||
|
||||
if conn != nil {
|
||||
logCodexWebsocketDisconnected(sessionID, authID, wsURL, reason, nil)
|
||||
if closer != nil {
|
||||
if errClose := closer.Close(); errClose != nil {
|
||||
log.Errorf("codex websockets executor: close websocket error: %v", errClose)
|
||||
}
|
||||
}
|
||||
}
|
||||
if lifecycle != nil {
|
||||
lifecycle.End(reason)
|
||||
}
|
||||
}
|
||||
|
||||
func logCodexWebsocketConnected(sessionID string, authID string, wsURL string) {
|
||||
log.Infof("codex websockets: upstream connected session=%s auth=%s url=%s", strings.TrimSpace(sessionID), strings.TrimSpace(authID), strings.TrimSpace(wsURL))
|
||||
}
|
||||
|
||||
func logCodexWebsocketDisconnected(sessionID string, authID string, wsURL string, reason string, err error) {
|
||||
if err != nil {
|
||||
log.Infof("codex websockets: upstream disconnected session=%s auth=%s url=%s reason=%s err=%v", strings.TrimSpace(sessionID), strings.TrimSpace(authID), strings.TrimSpace(wsURL), strings.TrimSpace(reason), err)
|
||||
return
|
||||
}
|
||||
log.Infof("codex websockets: upstream disconnected session=%s auth=%s url=%s reason=%s", strings.TrimSpace(sessionID), strings.TrimSpace(authID), strings.TrimSpace(wsURL), strings.TrimSpace(reason))
|
||||
}
|
||||
|
||||
// CloseCodexWebsocketSessionsForAuthID closes all active Codex upstream websocket sessions
|
||||
// associated with the supplied auth ID.
|
||||
func CloseCodexWebsocketSessionsForAuthID(authID string, reason string) {
|
||||
authID = strings.TrimSpace(authID)
|
||||
if authID == "" {
|
||||
return
|
||||
}
|
||||
reason = strings.TrimSpace(reason)
|
||||
if reason == "" {
|
||||
reason = "auth_removed"
|
||||
}
|
||||
|
||||
store := globalCodexWebsocketSessionStore
|
||||
if store == nil {
|
||||
return
|
||||
}
|
||||
|
||||
type sessionItem struct {
|
||||
sessionID string
|
||||
sess *codexWebsocketSession
|
||||
}
|
||||
|
||||
store.mu.Lock()
|
||||
items := make([]sessionItem, 0, len(store.sessions))
|
||||
for sessionID, sess := range store.sessions {
|
||||
items = append(items, sessionItem{sessionID: sessionID, sess: sess})
|
||||
}
|
||||
store.mu.Unlock()
|
||||
|
||||
matches := make([]sessionItem, 0)
|
||||
for i := range items {
|
||||
sess := items[i].sess
|
||||
if sess == nil {
|
||||
continue
|
||||
}
|
||||
sess.connMu.Lock()
|
||||
sessAuthID := strings.TrimSpace(sess.authID)
|
||||
sess.connMu.Unlock()
|
||||
if sessAuthID == authID {
|
||||
matches = append(matches, items[i])
|
||||
}
|
||||
}
|
||||
if len(matches) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
toClose := make([]*codexWebsocketSession, 0, len(matches))
|
||||
store.mu.Lock()
|
||||
for i := range matches {
|
||||
current, ok := store.sessions[matches[i].sessionID]
|
||||
if !ok || current == nil || current != matches[i].sess {
|
||||
continue
|
||||
}
|
||||
delete(store.sessions, matches[i].sessionID)
|
||||
toClose = append(toClose, current)
|
||||
}
|
||||
store.mu.Unlock()
|
||||
|
||||
for i := range toClose {
|
||||
closeCodexWebsocketSession(toClose[i], reason)
|
||||
}
|
||||
}
|
||||
Loading…
Reference in a new issue