273 lines
6.6 KiB
Go
273 lines
6.6 KiB
Go
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
|