diff --git a/internal/api/handlers/management/auth_files.go b/internal/api/handlers/management/auth_files.go index 5c044e0c..ecbdf851 100644 --- a/internal/api/handlers/management/auth_files.go +++ b/internal/api/handlers/management/auth_files.go @@ -2627,8 +2627,12 @@ func (h *Handler) GetAuthStatus(c *gin.Context) { return } - provider, status, isPlugin, metadata, ok := GetOAuthSessionDetails(state) + provider, status, isPlugin, metadata, completed, ok := GetOAuthSessionDetails(state) if !ok { + c.JSON(http.StatusOK, gin.H{"status": "error", "error": "unknown or expired state"}) + return + } + if completed { c.JSON(http.StatusOK, gin.H{"status": "ok"}) return } diff --git a/internal/api/handlers/management/oauth_callback.go b/internal/api/handlers/management/oauth_callback.go index 832462db..b0d3e9d5 100644 --- a/internal/api/handlers/management/oauth_callback.go +++ b/internal/api/handlers/management/oauth_callback.go @@ -86,11 +86,15 @@ func (h *Handler) handleOAuthCallback(c *gin.Context, req oauthCallbackRequest) return } - sessionProvider, sessionStatus, isPlugin, _, ok := GetOAuthSessionDetails(state) + sessionProvider, sessionStatus, isPlugin, _, completed, ok := GetOAuthSessionDetails(state) if !ok { c.JSON(http.StatusNotFound, gin.H{"status": "error", "error": "unknown or expired state"}) return } + if completed { + c.JSON(http.StatusConflict, gin.H{"status": "error", "error": "oauth flow is already completed"}) + return + } provider := strings.TrimSpace(req.Provider) if provider == "" { provider = sessionProvider diff --git a/internal/api/handlers/management/oauth_sessions.go b/internal/api/handlers/management/oauth_sessions.go index 078c51c6..a4318ffd 100644 --- a/internal/api/handlers/management/oauth_sessions.go +++ b/internal/api/handlers/management/oauth_sessions.go @@ -12,8 +12,9 @@ import ( ) const ( - oauthSessionTTL = 10 * time.Minute - maxOAuthStateLength = 128 + oauthSessionTTL = 10 * time.Minute + oauthCompletedSessionTTL = time.Minute + maxOAuthStateLength = 128 ) const ( @@ -33,23 +34,30 @@ type oauthSession struct { Status string Source string Metadata map[string]any + Completed bool CreatedAt time.Time ExpiresAt time.Time } type oauthSessionStore struct { - mu sync.RWMutex - ttl time.Duration - sessions map[string]oauthSession + mu sync.RWMutex + ttl time.Duration + completedTTL time.Duration + sessions map[string]oauthSession } func newOAuthSessionStore(ttl time.Duration) *oauthSessionStore { if ttl <= 0 { ttl = oauthSessionTTL } + completedTTL := oauthCompletedSessionTTL + if ttl < completedTTL { + completedTTL = ttl + } return &oauthSessionStore{ - ttl: ttl, - sessions: make(map[string]oauthSession), + ttl: ttl, + completedTTL: completedTTL, + sessions: make(map[string]oauthSession), } } @@ -127,7 +135,7 @@ func (s *oauthSessionStore) SetError(state, message string) { s.purgeExpiredLocked(now) session, ok := s.sessions[state] - if !ok { + if !ok || session.Completed { return } session.Status = message @@ -146,7 +154,15 @@ func (s *oauthSessionStore) Complete(state string) { defer s.mu.Unlock() s.purgeExpiredLocked(now) - delete(s.sessions, state) + session, ok := s.sessions[state] + if !ok { + return + } + session.Status = "" + session.Metadata = nil + session.Completed = true + session.ExpiresAt = now.Add(s.completedTTL) + s.sessions[state] = session } func (s *oauthSessionStore) CompleteProvider(provider string, source string) int { @@ -164,7 +180,11 @@ func (s *oauthSessionStore) CompleteProvider(provider string, source string) int removed := 0 for state, session := range s.sessions { if strings.EqualFold(session.Provider, provider) && (source == "" || session.Source == source) { - delete(s.sessions, state) + session.Status = "" + session.Metadata = nil + session.Completed = true + session.ExpiresAt = now.Add(s.completedTTL) + s.sessions[state] = session removed++ } } @@ -197,7 +217,7 @@ func (s *oauthSessionStore) IsPending(state, provider string) bool { if !ok { return false } - if session.Status != "" { + if session.Completed || session.Status != "" { return false } if provider == "" { @@ -245,12 +265,12 @@ func GetOAuthSession(state string) (provider string, status string, ok bool) { return session.Provider, session.Status, true } -func GetOAuthSessionDetails(state string) (provider string, status string, isPlugin bool, metadata map[string]any, ok bool) { +func GetOAuthSessionDetails(state string) (provider string, status string, isPlugin bool, metadata map[string]any, completed bool, ok bool) { session, ok := oauthSessions.Get(state) if !ok { - return "", "", false, nil, false + return "", "", false, nil, false, false } - return session.Provider, session.Status, session.Source == oauthSessionSourcePlugin, cloneOAuthSessionMetadata(session.Metadata), true + return session.Provider, session.Status, session.Source == oauthSessionSourcePlugin, cloneOAuthSessionMetadata(session.Metadata), session.Completed, true } func IsOAuthSessionPending(state, provider string) bool { diff --git a/internal/api/handlers/management/oauth_sessions_test.go b/internal/api/handlers/management/oauth_sessions_test.go new file mode 100644 index 00000000..8e163d35 --- /dev/null +++ b/internal/api/handlers/management/oauth_sessions_test.go @@ -0,0 +1,102 @@ +package management + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/router-for-me/CLIProxyAPI/v7/internal/config" +) + +func TestOAuthSessionStoreCompleteKeepsShortLivedSession(t *testing.T) { + store := newOAuthSessionStore(time.Minute) + store.Register("completed-state", "codex") + + store.Complete("completed-state") + + if _, ok := store.Get("completed-state"); !ok { + t.Fatal("completed OAuth session was deleted instead of retained as a tombstone") + } + if store.IsPending("completed-state", "codex") { + t.Fatal("completed OAuth session remained pending") + } +} + +func TestGetAuthStatusRejectsUnknownStateAndAcceptsCompletedState(t *testing.T) { + store := newOAuthSessionStore(time.Minute) + replaceOAuthSessionStoreForTest(t, store) + + handler := &Handler{} + router := gin.New() + router.GET("/status", handler.GetAuthStatus) + + unknown := performOAuthStatusRequest(t, router, "unknown-state") + if unknown.Status != "error" || unknown.Error != "unknown or expired state" { + t.Fatalf("unknown state response = %#v, want unknown/expired error", unknown) + } + + store.Register("completed-state", "codex") + store.Complete("completed-state") + completed := performOAuthStatusRequest(t, router, "completed-state") + if completed.Status != "ok" || completed.Error != "" { + t.Fatalf("completed state response = %#v, want success", completed) + } +} + +func TestOAuthCallbackRejectsCompletedSession(t *testing.T) { + store := newOAuthSessionStore(time.Minute) + replaceOAuthSessionStoreForTest(t, store) + store.Register("completed-state", "codex") + store.Complete("completed-state") + + handler := NewHandlerWithoutConfigFilePath(&config.Config{AuthDir: t.TempDir()}, nil) + router := gin.New() + router.POST("/oauth-callback", handler.PostOAuthCallback) + + req := httptest.NewRequest( + http.MethodPost, + "/oauth-callback", + strings.NewReader(`{"provider":"codex","state":"completed-state","code":"test-code"}`), + ) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + + if w.Code != http.StatusConflict { + t.Fatalf("completed callback status = %d, want %d; body=%s", w.Code, http.StatusConflict, w.Body.String()) + } +} + +type oauthStatusResponse struct { + Status string `json:"status"` + Error string `json:"error"` +} + +func performOAuthStatusRequest(t *testing.T, router http.Handler, state string) oauthStatusResponse { + t.Helper() + req := httptest.NewRequest(http.MethodGet, "/status?state="+state, nil) + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + + if w.Code != http.StatusOK { + t.Fatalf("status request returned %d, want %d; body=%s", w.Code, http.StatusOK, w.Body.String()) + } + var response oauthStatusResponse + if errDecode := json.Unmarshal(w.Body.Bytes(), &response); errDecode != nil { + t.Fatalf("decode status response: %v", errDecode) + } + return response +} + +func replaceOAuthSessionStoreForTest(t *testing.T, store *oauthSessionStore) { + t.Helper() + original := oauthSessions + oauthSessions = store + t.Cleanup(func() { + oauthSessions = original + }) +}