Files
codex-helper/backend/internal/app/app.go
T

370 lines
9.3 KiB
Go

package app
import (
"context"
"crypto/tls"
"database/sql"
"embed"
"encoding/json"
"errors"
"io"
"io/fs"
"log"
"net/http"
"net/smtp"
"os"
"path/filepath"
"strconv"
"strings"
"sync"
"time"
"codex-helper/internal/codex"
"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.2.0"
type App struct {
dataDir string
store *store.Store
vault *security.Vault
server *http.Server
ctx context.Context
cancel context.CancelFunc
mu sync.RWMutex
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
}
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
}
v, e := security.OpenVault(dir)
if e != nil {
return nil, e
}
ctx, cancel := context.WithCancel(context.Background())
a := &App{dataDir: dir, store: s, vault: v, 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 {
if e := os.MkdirAll(filepath.Join(a.dataDir, "codex"), 0700); e != nil {
return e
}
go a.keepCodex()
go a.scheduler()
go a.telegramLoop()
log.Printf("codex-helper listening on %s", a.server.Addr)
e := a.server.ListenAndServe()
if errors.Is(e, http.ErrServerClosed) {
return nil
}
return e
}
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()
}
a.mu.RUnlock()
_ = 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")
}
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()
return a.runtimes[id]
}
func (a *App) syncAll(ctx context.Context) {
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 {
_ = a.syncAccount(ctx, id)
}
}
func (a *App) scheduler() {
t := time.NewTicker(time.Minute)
defer t.Stop()
for {
select {
case <-a.ctx.Done():
return
case <-t.C:
g := a.general()
if time.Now().Unix()%(int64(g.SyncMinutes)*60) < 60 {
a.syncAll(context.Background())
}
_, _ = a.store.Cleanup(g.RetentionDays)
go a.processReminders()
}
}
}
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/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()})
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})
})
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 {
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, e := r.Cookie("session")
if e != 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": "未登录"})
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
}
}
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