package app import ( "context" "embed" "encoding/json" "errors" "io" "io/fs" "log" "net/http" "os" "strings" "sync" "time" "codex-helper/internal/cliproxy" "codex-helper/internal/security" "codex-helper/internal/store" webassets "codex-helper/internal/web" ) // Version is overridden at build time for release images. var Version = "0.3.0" 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 { dash Dashboard syncing sync.Mutex } type cpaClient interface { Configured() bool Snapshot(context.Context, string) (cliproxy.Snapshot, error) } const automaticSyncInterval = time.Duration(automaticSyncMinutes) * time.Minute func New() (*App, error) { dir := env("DATA_DIR", "/data") s, err := store.Open(dir) if err != nil { return nil, err } 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, 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) } 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 { go a.syncScheduler() go a.maintenanceScheduler() go a.telegramLoop() log.Printf("codex-helper listening on %s", a.server.Addr) err := a.server.ListenAndServe() if errors.Is(err, http.ErrServerClosed) { return nil } return err } func (a *App) Close() { a.cancel() ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() if a.server != nil { _ = a.server.Shutdown(ctx) } a.syncAllMu.Lock() a.syncAllMu.Unlock() _ = a.store.DB.Close() } func (a *App) addRuntime(id int64) { 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} } a.mu.Lock() a.runtimes[id] = &accountRuntime{dash: dashboard} a.mu.Unlock() } 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 { 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) syncScheduler() { a.syncAll(a.ctx) ticker := time.NewTicker(automaticSyncInterval) defer ticker.Stop() for { select { case <-a.ctx.Done(): return 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() } } } 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, http.StatusOK, map[string]any{"status": "ok"}) }) m.HandleFunc("/health/ready", func(w http.ResponseWriter, r *http.Request) { if err := a.store.Health(r.Context()); err != nil { jsonOut(w, http.StatusServiceUnavailable, map[string]string{"error": err.Error()}) return } 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 _, err := fs.Stat(sub, strings.TrimPrefix(r.URL.Path, "/")); err == nil { files.ServeHTTP(w, r) return } } b, _ := fs.ReadFile(sub, "index.html") w.Header().Set("Content-Type", "text/html; charset=utf-8") _, _ = w.Write(b) }) 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") w.Header().Set("X-Frame-Options", "DENY") w.Header().Set("Referrer-Policy", "same-origin") w.Header().Set("Content-Security-Policy", "default-src 'self'; style-src 'self' 'unsafe-inline'; img-src 'self' data:; connect-src 'self'") 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, 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, http.StatusUnauthorized, 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