Fix account service startup race

This commit is contained in:
zhoujun0601
2026-08-13 07:33:04 -04:00
parent 4c05dde0f1
commit d6edde6534
5 changed files with 259 additions and 29 deletions
+14 -8
View File
@@ -21,7 +21,7 @@ func (a *App) api(w http.ResponseWriter, r *http.Request) {
connected := false
a.mu.RLock()
for _, rt := range a.runtimes {
if rt.client.Connected() {
if rt.Ready() {
connected = true
break
}
@@ -178,7 +178,7 @@ func (a *App) accountAPI(w http.ResponseWriter, r *http.Request, p string) {
a.mu.Lock()
delete(a.runtimes, id)
a.mu.Unlock()
_ = rt.client.Close()
rt.stop()
e = a.store.DeleteAccount(id)
if e == nil {
dir := filepath.Join(a.dataDir, "accounts", strconv.FormatInt(id, 10))
@@ -194,7 +194,10 @@ func (a *App) accountAPI(w http.ResponseWriter, r *http.Request, p string) {
a.deviceLogin(w, r, id)
case action == "logout" && r.Method == "POST":
var out any
e = rt.client.Call(r.Context(), "account/logout", map[string]any{}, &out)
e = rt.ensureReady(r.Context())
if e == nil {
e = rt.client.Call(r.Context(), "account/logout", map[string]any{}, &out)
}
if e == nil {
_ = a.store.UpdateAccount(id, nil, nil, false)
jsonOut(w, 200, map[string]bool{"ok": true})
@@ -326,11 +329,14 @@ func (a *App) generalAPI(w http.ResponseWriter, r *http.Request) {
func (a *App) deviceLogin(w http.ResponseWriter, r *http.Request, id int64) {
var out map[string]any
rt := a.runtime(id)
if !rt.client.Connected() {
jsonOut(w, 503, map[string]string{"error": "该账号服务正在启动,请稍后重试"})
if rt == nil {
jsonOut(w, 404, map[string]string{"error": "账号不存在"})
return
}
e := rt.client.Call(r.Context(), "account/login/start", map[string]any{"type": "chatgptDeviceCode"}, &out)
e := rt.ensureReady(r.Context())
if e == nil {
e = rt.client.Call(r.Context(), "account/login/start", map[string]any{"type": "chatgptDeviceCode"}, &out)
}
if e != nil {
jsonOut(w, 502, map[string]string{"error": e.Error()})
return
@@ -345,8 +351,8 @@ func (a *App) syncAccount(ctx context.Context, id int64) error {
}
rt.syncing.Lock()
defer rt.syncing.Unlock()
if !rt.client.Connected() {
return fmt.Errorf("app-server 未连接")
if err := rt.ensureReady(ctx); err != nil {
return err
}
ctx, c := context.WithTimeout(ctx, 20*time.Second)
defer c()
+83 -19
View File
@@ -37,9 +37,22 @@ type App struct {
loginAttempts sync.Map
}
type accountRuntime struct {
client *codex.Client
dash Dashboard
syncing sync.Mutex
client codexClient
processCtx context.Context
dash Dashboard
syncing sync.Mutex
lifecycle sync.Mutex
stateMu sync.RWMutex
ready bool
stopped bool
}
type codexClient interface {
Start(context.Context) error
Initialize(context.Context) error
Call(context.Context, string, any, any) error
Close() error
Connected() bool
}
func New() (*App, error) {
@@ -88,7 +101,7 @@ func (a *App) Close() {
_ = a.server.Shutdown(ctx)
a.mu.RLock()
for _, rt := range a.runtimes {
_ = rt.client.Close()
rt.stop()
}
a.mu.RUnlock()
_ = a.store.DB.Close()
@@ -108,21 +121,13 @@ func (a *App) keepCodex() {
a.mu.RUnlock()
for _, id := range ids {
rt := a.runtime(id)
if rt == nil || rt.client.Connected() {
if rt == nil || rt.Ready() {
continue
}
if e := rt.client.Start(a.ctx); e == nil {
ctx, c := context.WithTimeout(a.ctx, 20*time.Second)
e = rt.client.Initialize(ctx)
c()
if e == nil {
_ = a.syncAccount(context.Background(), id)
} else {
log.Printf("app-server initialize: %v", e)
// A live process is not necessarily an initialized process. Tear it
// down so the next iteration starts a fresh protocol session.
_ = rt.client.Close()
}
if e := rt.ensureReady(a.ctx); e == nil {
_ = a.syncAccount(context.Background(), id)
} else if !errors.Is(e, errRuntimeStopped) {
log.Printf("app-server initialize: %v", e)
}
}
time.Sleep(time.Second)
@@ -142,11 +147,70 @@ func (a *App) addRuntime(id int64) {
// existing login, and fresh installs use the same deterministic path.
dir = filepath.Join(a.dataDir, "codex")
}
rt := &accountRuntime{client: codex.New(dir, a.onCodexNotification(id)), dash: Dashboard{Limits: []LimitBucket{}, Usage: []UsagePoint{}, Stale: true}}
rt := &accountRuntime{client: codex.New(dir, a.onCodexNotification(id)), processCtx: a.ctx, dash: Dashboard{Limits: []LimitBucket{}, Usage: []UsagePoint{}, Stale: true}}
a.mu.Lock()
a.runtimes[id] = rt
a.mu.Unlock()
}
var errRuntimeStopped = errors.New("账号服务已停止")
func (rt *accountRuntime) ensureReady(ctx context.Context) error {
rt.lifecycle.Lock()
defer rt.lifecycle.Unlock()
rt.stateMu.RLock()
stopped, ready := rt.stopped, rt.ready
rt.stateMu.RUnlock()
if stopped {
return errRuntimeStopped
}
if ready && rt.client.Connected() {
return nil
}
// Connected only means the child process is alive. If a previous attempt
// did not finish the protocol handshake, discard it before trying again.
rt.stateMu.Lock()
rt.ready = false
rt.stateMu.Unlock()
if rt.client.Connected() {
_ = rt.client.Close()
}
processCtx := rt.processCtx
if processCtx == nil {
processCtx = ctx
}
if err := rt.client.Start(processCtx); err != nil {
return err
}
initCtx, cancel := context.WithTimeout(ctx, 20*time.Second)
err := rt.client.Initialize(initCtx)
cancel()
if err != nil {
_ = rt.client.Close()
return err
}
rt.stateMu.Lock()
rt.ready = true
rt.stateMu.Unlock()
return nil
}
func (rt *accountRuntime) Ready() bool {
rt.stateMu.RLock()
ready := !rt.stopped && rt.ready
rt.stateMu.RUnlock()
return ready && rt.client.Connected()
}
func (rt *accountRuntime) stop() {
rt.lifecycle.Lock()
defer rt.lifecycle.Unlock()
rt.stateMu.Lock()
rt.stopped = true
rt.ready = false
rt.stateMu.Unlock()
_ = rt.client.Close()
}
func (a *App) runtime(id int64) *accountRuntime {
a.mu.RLock()
defer a.mu.RUnlock()
@@ -192,7 +256,7 @@ func (a *App) routes() http.Handler {
connected := false
a.mu.RLock()
for _, rt := range a.runtimes {
if rt.client.Connected() {
if rt.Ready() {
connected = true
break
}
+160
View File
@@ -0,0 +1,160 @@
package app
import (
"context"
"errors"
"net/http/httptest"
"sync"
"testing"
)
type fakeCodexClient struct {
mu sync.Mutex
connected bool
starts int
initializes int
closes int
calls int
initErrors []error
initStarted chan struct{}
initRelease chan struct{}
}
func (f *fakeCodexClient) Start(context.Context) error {
f.mu.Lock()
defer f.mu.Unlock()
f.starts++
f.connected = true
return nil
}
func (f *fakeCodexClient) Initialize(context.Context) error {
f.mu.Lock()
f.initializes++
var err error
if len(f.initErrors) > 0 {
err, f.initErrors = f.initErrors[0], f.initErrors[1:]
}
started, release := f.initStarted, f.initRelease
f.mu.Unlock()
if started != nil {
select {
case started <- struct{}{}:
default:
}
}
if release != nil {
<-release
}
return err
}
func (f *fakeCodexClient) Call(_ context.Context, method string, _ any, out any) error {
f.mu.Lock()
defer f.mu.Unlock()
f.calls++
if method == "account/login/start" {
result := out.(*map[string]any)
*result = map[string]any{"verificationUrl": "https://example.test/device", "userCode": "ABCD-EFGH"}
}
return nil
}
func (f *fakeCodexClient) Close() error {
f.mu.Lock()
defer f.mu.Unlock()
f.closes++
f.connected = false
return nil
}
func (f *fakeCodexClient) Connected() bool {
f.mu.Lock()
defer f.mu.Unlock()
return f.connected
}
func (f *fakeCodexClient) counts() (starts, initializes, closes, calls int) {
f.mu.Lock()
defer f.mu.Unlock()
return f.starts, f.initializes, f.closes, f.calls
}
func TestEnsureReadySerializesColdStart(t *testing.T) {
client := &fakeCodexClient{}
rt := &accountRuntime{client: client}
var wg sync.WaitGroup
errs := make(chan error, 8)
for range 8 {
wg.Add(1)
go func() {
defer wg.Done()
errs <- rt.ensureReady(context.Background())
}()
}
wg.Wait()
close(errs)
for err := range errs {
if err != nil {
t.Fatal(err)
}
}
starts, initializes, _, _ := client.counts()
if starts != 1 || initializes != 1 {
t.Fatalf("cold starts = %d, initializes = %d; want 1 each", starts, initializes)
}
}
func TestEnsureReadyRetriesAfterInitializeFailure(t *testing.T) {
client := &fakeCodexClient{initErrors: []error{errors.New("handshake failed")}}
rt := &accountRuntime{client: client}
if err := rt.ensureReady(context.Background()); err == nil {
t.Fatal("first initialization unexpectedly succeeded")
}
if err := rt.ensureReady(context.Background()); err != nil {
t.Fatalf("retry failed: %v", err)
}
starts, initializes, closes, _ := client.counts()
if starts != 2 || initializes != 2 || closes != 1 {
t.Fatalf("starts = %d, initializes = %d, closes = %d; want 2, 2, 1", starts, initializes, closes)
}
}
func TestStopWaitsForStartupAndPreventsRestart(t *testing.T) {
started := make(chan struct{}, 1)
release := make(chan struct{})
client := &fakeCodexClient{initStarted: started, initRelease: release}
rt := &accountRuntime{client: client}
readyDone := make(chan error, 1)
go func() { readyDone <- rt.ensureReady(context.Background()) }()
<-started
stopDone := make(chan struct{})
go func() { rt.stop(); close(stopDone) }()
close(release)
if err := <-readyDone; err != nil {
t.Fatalf("startup failed: %v", err)
}
<-stopDone
if err := rt.ensureReady(context.Background()); !errors.Is(err, errRuntimeStopped) {
t.Fatalf("restart error = %v; want stopped", err)
}
starts, initializes, closes, _ := client.counts()
if starts != 1 || initializes != 1 || closes != 1 {
t.Fatalf("starts = %d, initializes = %d, closes = %d; want 1 each", starts, initializes, closes)
}
}
func TestDeviceLoginStartsColdRuntime(t *testing.T) {
client := &fakeCodexClient{}
a := &App{runtimes: map[int64]*accountRuntime{2: {client: client}}}
recorder := httptest.NewRecorder()
request := httptest.NewRequest("POST", "/api/v1/accounts/2/login/device", nil)
a.deviceLogin(recorder, request, 2)
if recorder.Code != 200 {
t.Fatalf("status = %d, body = %s", recorder.Code, recorder.Body.String())
}
starts, initializes, _, calls := client.counts()
if starts != 1 || initializes != 1 || calls != 1 {
t.Fatalf("starts = %d, initializes = %d, calls = %d; want 1 each", starts, initializes, calls)
}
}
+1 -1
View File
@@ -1,3 +1,3 @@
<script type="module" crossorigin src="/assets/index-CCkNuko2.js"></script>
<script type="module" crossorigin src="/assets/index-D5kXUJ5m.js"></script>
<link rel="stylesheet" crossorigin href="/assets/index-2I4Pedo2.css">
<div id="root"></div>