This commit is contained in:
+525
-376
File diff suppressed because it is too large
Load Diff
+111
-211
@@ -2,8 +2,6 @@ package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"database/sql"
|
||||
"embed"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
@@ -11,15 +9,12 @@ import (
|
||||
"io/fs"
|
||||
"log"
|
||||
"net/http"
|
||||
"net/smtp"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"codex-helper/internal/codex"
|
||||
"codex-helper/internal/cliproxy"
|
||||
"codex-helper/internal/security"
|
||||
"codex-helper/internal/store"
|
||||
webassets "codex-helper/internal/web"
|
||||
@@ -32,48 +27,51 @@ type App struct {
|
||||
dataDir string
|
||||
store *store.Store
|
||||
vault *security.Vault
|
||||
cpa cpaClient
|
||||
server *http.Server
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
mu sync.RWMutex
|
||||
syncAllMu sync.Mutex
|
||||
runtimes map[int64]*accountRuntime
|
||||
loginAttempts sync.Map
|
||||
reminderMu sync.Mutex
|
||||
telegramMu sync.Mutex
|
||||
}
|
||||
|
||||
type accountRuntime struct {
|
||||
client codexClient
|
||||
processCtx context.Context
|
||||
dash Dashboard
|
||||
syncing sync.Mutex
|
||||
lifecycle sync.Mutex
|
||||
stateMu sync.RWMutex
|
||||
ready bool
|
||||
stopped bool
|
||||
dash Dashboard
|
||||
syncing sync.Mutex
|
||||
}
|
||||
|
||||
type cpaClient interface {
|
||||
Configured() bool
|
||||
Snapshot(context.Context, string) (cliproxy.Snapshot, error)
|
||||
}
|
||||
|
||||
const automaticSyncInterval = time.Duration(automaticSyncMinutes) * time.Minute
|
||||
|
||||
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) {
|
||||
dir := env("DATA_DIR", "/data")
|
||||
s, e := store.Open(dir)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
s, err := store.Open(dir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
v, e := security.OpenVault(dir)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
v, err := security.OpenVault(dir)
|
||||
if err != nil {
|
||||
_ = s.DB.Close()
|
||||
return nil, err
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
a := &App{dataDir: dir, store: s, vault: v, ctx: ctx, cancel: cancel, runtimes: map[int64]*accountRuntime{}}
|
||||
a := &App{
|
||||
dataDir: dir,
|
||||
store: s,
|
||||
vault: v,
|
||||
cpa: cliproxy.New(os.Getenv("CLIPROXY_API_BASE_URL"), os.Getenv("CLIPROXY_API_MANAGEMENT_KEY")),
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
runtimes: map[int64]*accountRuntime{},
|
||||
}
|
||||
accounts, _ := s.Accounts()
|
||||
for _, account := range accounts {
|
||||
a.addRuntime(account.ID)
|
||||
@@ -81,205 +79,111 @@ func New() (*App, error) {
|
||||
a.server = &http.Server{Addr: env("LISTEN_ADDR", ":8080"), Handler: a.routes(), ReadHeaderTimeout: 10 * time.Second, IdleTimeout: 60 * time.Second}
|
||||
return a, nil
|
||||
}
|
||||
|
||||
func env(k, d string) string {
|
||||
if v := os.Getenv(k); v != "" {
|
||||
return v
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
func (a *App) Run() error {
|
||||
if e := os.MkdirAll(filepath.Join(a.dataDir, "codex"), 0700); e != nil {
|
||||
return e
|
||||
}
|
||||
go a.keepCodex()
|
||||
go a.scheduler()
|
||||
go a.syncScheduler()
|
||||
go a.maintenanceScheduler()
|
||||
go a.telegramLoop()
|
||||
log.Printf("codex-helper listening on %s", a.server.Addr)
|
||||
e := a.server.ListenAndServe()
|
||||
if errors.Is(e, http.ErrServerClosed) {
|
||||
err := a.server.ListenAndServe()
|
||||
if errors.Is(err, http.ErrServerClosed) {
|
||||
return nil
|
||||
}
|
||||
return e
|
||||
return err
|
||||
}
|
||||
|
||||
func (a *App) Close() {
|
||||
a.cancel()
|
||||
ctx, c := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer c()
|
||||
_ = a.server.Shutdown(ctx)
|
||||
a.mu.RLock()
|
||||
for _, rt := range a.runtimes {
|
||||
rt.stop()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
if a.server != nil {
|
||||
_ = a.server.Shutdown(ctx)
|
||||
}
|
||||
a.mu.RUnlock()
|
||||
a.syncAllMu.Lock()
|
||||
a.syncAllMu.Unlock()
|
||||
_ = a.store.DB.Close()
|
||||
}
|
||||
func (a *App) keepCodex() {
|
||||
for {
|
||||
select {
|
||||
case <-a.ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
a.mu.RLock()
|
||||
ids := make([]int64, 0, len(a.runtimes))
|
||||
for id := range a.runtimes {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
a.mu.RUnlock()
|
||||
for _, id := range ids {
|
||||
rt := a.runtime(id)
|
||||
if rt == nil || rt.Ready() {
|
||||
continue
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
func (a *App) onCodexNotification(id int64) func(string, json.RawMessage) {
|
||||
return func(method string, _ json.RawMessage) {
|
||||
if method == "account/login/completed" || method == "account/updated" || method == "account/rateLimits/updated" {
|
||||
go a.syncAccountWithRetry(id, method == "account/login/completed")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) syncAccountWithRetry(id int64, requireClassifiedAccount bool) {
|
||||
delays := []time.Duration{0, time.Second, 3 * time.Second}
|
||||
var err error
|
||||
for _, delay := range delays {
|
||||
if delay > 0 {
|
||||
select {
|
||||
case <-a.ctx.Done():
|
||||
return
|
||||
case <-time.After(delay):
|
||||
}
|
||||
}
|
||||
err = a.syncAccount(context.Background(), id)
|
||||
if err == nil {
|
||||
rt := a.runtime(id)
|
||||
if !requireClassifiedAccount || rt == nil || accountClassificationReady(rt) {
|
||||
return
|
||||
}
|
||||
err = errors.New("登录已完成,但工作区套餐尚未就绪")
|
||||
}
|
||||
}
|
||||
if rt := a.runtime(id); rt != nil {
|
||||
rt.syncing.Lock()
|
||||
rt.dash.Stale = true
|
||||
rt.dash.LastError = err.Error()
|
||||
rt.syncing.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
func accountClassificationReady(rt *accountRuntime) bool {
|
||||
rt.syncing.Lock()
|
||||
defer rt.syncing.Unlock()
|
||||
return rt.dash.Account.Connected && store.AccountKind(rt.dash.Account.PlanType) != "unknown"
|
||||
}
|
||||
func (a *App) addRuntime(id int64) {
|
||||
dir := filepath.Join(a.dataDir, "accounts", strconv.FormatInt(id, 10), "codex")
|
||||
if id == 1 {
|
||||
// Account 1 deliberately keeps the legacy path so upgrades retain the
|
||||
// existing login, and fresh installs use the same deterministic path.
|
||||
dir = filepath.Join(a.dataDir, "codex")
|
||||
dashboard := Dashboard{AccountID: id, Limits: []LimitBucket{}, Usage: []UsagePoint{}, Stale: true}
|
||||
if account, err := a.store.Account(id); err == nil {
|
||||
dashboard.DisplayName = account.DisplayName
|
||||
dashboard.Account = AccountView{Email: account.Email, PlanType: account.PlanType, Connected: account.Connected}
|
||||
}
|
||||
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.runtimes[id] = &accountRuntime{dash: dashboard}
|
||||
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()
|
||||
return a.runtimes[id]
|
||||
}
|
||||
|
||||
func (a *App) syncAll(ctx context.Context) {
|
||||
if !a.syncAllMu.TryLock() {
|
||||
return
|
||||
}
|
||||
defer a.syncAllMu.Unlock()
|
||||
if ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
a.mu.RLock()
|
||||
ids := make([]int64, 0, len(a.runtimes))
|
||||
for id := range a.runtimes {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
a.mu.RUnlock()
|
||||
|
||||
const maxConcurrentAccountSyncs = 4
|
||||
sem := make(chan struct{}, maxConcurrentAccountSyncs)
|
||||
var wg sync.WaitGroup
|
||||
for _, id := range ids {
|
||||
_ = a.syncAccount(ctx, id)
|
||||
wg.Add(1)
|
||||
go func(accountID int64) {
|
||||
defer wg.Done()
|
||||
select {
|
||||
case sem <- struct{}{}:
|
||||
defer func() { <-sem }()
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
_ = a.syncAccount(ctx, accountID)
|
||||
}(id)
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
func (a *App) scheduler() {
|
||||
maintenanceTicker := time.NewTicker(time.Minute)
|
||||
syncTicker := time.NewTicker(automaticSyncInterval)
|
||||
defer maintenanceTicker.Stop()
|
||||
defer syncTicker.Stop()
|
||||
|
||||
func (a *App) syncScheduler() {
|
||||
a.syncAll(a.ctx)
|
||||
ticker := time.NewTicker(automaticSyncInterval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-a.ctx.Done():
|
||||
return
|
||||
case <-syncTicker.C:
|
||||
a.syncAll(context.Background())
|
||||
case <-maintenanceTicker.C:
|
||||
case <-ticker.C:
|
||||
a.syncAll(a.ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) maintenanceScheduler() {
|
||||
ticker := time.NewTicker(time.Minute)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-a.ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
g := a.general()
|
||||
_, _ = a.store.Cleanup(g.RetentionDays)
|
||||
go a.processReminders()
|
||||
@@ -287,31 +191,28 @@ func (a *App) scheduler() {
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) cpaConfigured() bool {
|
||||
return a.cpa != nil && a.cpa.Configured()
|
||||
}
|
||||
|
||||
func (a *App) routes() http.Handler {
|
||||
m := http.NewServeMux()
|
||||
m.HandleFunc("/health/live", func(w http.ResponseWriter, r *http.Request) { jsonOut(w, 200, map[string]any{"status": "ok"}) })
|
||||
m.HandleFunc("/health/live", func(w http.ResponseWriter, r *http.Request) {
|
||||
jsonOut(w, http.StatusOK, map[string]any{"status": "ok"})
|
||||
})
|
||||
m.HandleFunc("/health/ready", func(w http.ResponseWriter, r *http.Request) {
|
||||
if e := a.store.Health(r.Context()); e != nil {
|
||||
jsonOut(w, 503, map[string]string{"error": e.Error()})
|
||||
if err := a.store.Health(r.Context()); err != nil {
|
||||
jsonOut(w, http.StatusServiceUnavailable, map[string]string{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
connected := false
|
||||
a.mu.RLock()
|
||||
for _, rt := range a.runtimes {
|
||||
if rt.Ready() {
|
||||
connected = true
|
||||
break
|
||||
}
|
||||
}
|
||||
a.mu.RUnlock()
|
||||
jsonOut(w, 200, map[string]any{"status": "ok", "appServer": connected})
|
||||
jsonOut(w, http.StatusOK, map[string]any{"status": "ok", "cpa": a.cpaConfigured()})
|
||||
})
|
||||
m.HandleFunc("/api/v1/", a.api)
|
||||
sub, _ := fs.Sub(webassets.Assets, "dist")
|
||||
files := http.FileServer(http.FS(sub))
|
||||
m.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/" {
|
||||
if _, e := fs.Stat(sub, strings.TrimPrefix(r.URL.Path, "/")); e == nil {
|
||||
if _, err := fs.Stat(sub, strings.TrimPrefix(r.URL.Path, "/")); err == nil {
|
||||
files.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
@@ -322,6 +223,7 @@ func (a *App) routes() http.Handler {
|
||||
})
|
||||
return securityHeaders(m)
|
||||
}
|
||||
|
||||
func securityHeaders(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("X-Content-Type-Options", "nosniff")
|
||||
@@ -331,42 +233,40 @@ func securityHeaders(next http.Handler) http.Handler {
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func jsonOut(w http.ResponseWriter, status int, v any) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(status)
|
||||
_ = json.NewEncoder(w).Encode(v)
|
||||
}
|
||||
|
||||
func decode(r *http.Request, v any) error {
|
||||
defer r.Body.Close()
|
||||
d := json.NewDecoder(io.LimitReader(r.Body, 1<<20))
|
||||
d.DisallowUnknownFields()
|
||||
return d.Decode(v)
|
||||
}
|
||||
|
||||
func (a *App) authed(r *http.Request) bool {
|
||||
c, e := r.Cookie("session")
|
||||
if e != nil {
|
||||
c, err := r.Cookie("session")
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
var x int
|
||||
return a.store.DB.QueryRow("SELECT 1 FROM sessions WHERE token_hash=? AND expires_at>?", security.HashToken(c.Value), time.Now().Unix()).Scan(&x) == nil
|
||||
}
|
||||
|
||||
func (a *App) require(w http.ResponseWriter, r *http.Request) bool {
|
||||
if !a.authed(r) {
|
||||
jsonOut(w, 401, map[string]string{"error": "未登录"})
|
||||
jsonOut(w, http.StatusUnauthorized, map[string]string{"error": "未登录"})
|
||||
return false
|
||||
}
|
||||
if r.Method != "GET" && r.Method != "HEAD" {
|
||||
if r.Header.Get("X-Requested-With") != "codex-helper" {
|
||||
jsonOut(w, 403, map[string]string{"error": "请求来源校验失败"})
|
||||
return false
|
||||
}
|
||||
if r.Method != http.MethodGet && r.Method != http.MethodHead && r.Header.Get("X-Requested-With") != "codex-helper" {
|
||||
jsonOut(w, http.StatusForbidden, map[string]string{"error": "请求来源校验失败"})
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// Keep an explicit embed reference visible to tooling.
|
||||
var _ embed.FS
|
||||
var _ = sql.ErrNoRows
|
||||
var _ = strconv.Itoa
|
||||
var _ = tls.VersionTLS13
|
||||
var _ = smtp.SendMail
|
||||
|
||||
@@ -415,7 +415,9 @@ func (a *App) handleTG(t TelegramSettings, chat int64, text string) {
|
||||
return
|
||||
}
|
||||
if text == "立即刷新" || text == "/refresh" {
|
||||
a.syncAll(context.Background())
|
||||
go a.syncAll(a.ctx)
|
||||
_ = tgSend(t, "🔄 <b>刷新请求已提交</b>\n\n如果已有刷新正在执行,将复用该轮结果;稍后可再次查询当前用量。")
|
||||
return
|
||||
}
|
||||
msg := ""
|
||||
a.mu.RLock()
|
||||
|
||||
@@ -367,6 +367,13 @@ func reminderDashboard(fetchedAt int64, used float64, resetsAt int64) Dashboard
|
||||
}
|
||||
}
|
||||
|
||||
func ensureReminderAccount(t *testing.T, a *App) {
|
||||
t.Helper()
|
||||
if _, err := a.store.CreateAccount("测试账号", "test-auth"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func notificationCount(t *testing.T, a *App) int {
|
||||
t.Helper()
|
||||
var count int
|
||||
@@ -378,6 +385,7 @@ func notificationCount(t *testing.T, a *App) int {
|
||||
|
||||
func TestStoreLimitSnapshotsDetectsEarlyReset(t *testing.T) {
|
||||
a := newReminderTestApp(t)
|
||||
ensureReminderAccount(t, a)
|
||||
now := time.Now().Unix()
|
||||
if detected, err := a.storeLimitSnapshots(reminderDashboard(now, 42, now+3600)); err != nil || detected {
|
||||
t.Fatalf("initial snapshot: detected=%v err=%v", detected, err)
|
||||
@@ -433,6 +441,7 @@ func TestLegacyNotificationFormatting(t *testing.T) {
|
||||
|
||||
func TestStoreLimitSnapshotsUsesScheduledAfterDedupeKey(t *testing.T) {
|
||||
a := newReminderTestApp(t)
|
||||
ensureReminderAccount(t, a)
|
||||
now := time.Now().Unix()
|
||||
resetAt := now + 30
|
||||
_, _ = a.storeLimitSnapshots(reminderDashboard(now, 70, resetAt))
|
||||
@@ -466,6 +475,7 @@ func TestStoreLimitSnapshotsIgnoresNonResetChanges(t *testing.T) {
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
a := newReminderTestApp(t)
|
||||
ensureReminderAccount(t, a)
|
||||
g := defaults()
|
||||
g.NotifyAfter = tt.notifyAfter
|
||||
if err := a.store.SetJSON("general", g); err != nil {
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user