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
@@ -0,0 +1,745 @@
|
||||
package cliproxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
usageURL = "https://chatgpt.com/backend-api/wham/usage"
|
||||
profileURL = "https://chatgpt.com/backend-api/wham/profiles/me"
|
||||
resetCreditsURL = "https://chatgpt.com/backend-api/wham/rate-limit-reset-credits"
|
||||
codexUserAgent = "codex_cli_rs/0.76.0 (Debian 13.0.0; x86_64) WindowsTerminal"
|
||||
maxResponseBody = 4 << 20
|
||||
)
|
||||
|
||||
var ErrNotConfigured = errors.New("CLIProxyAPI Management API 未配置")
|
||||
|
||||
type Client struct {
|
||||
baseURL string
|
||||
managementKey string
|
||||
httpClient *http.Client
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
type AuthFile struct {
|
||||
AuthIndex string
|
||||
Label string
|
||||
Name string
|
||||
Email *string
|
||||
Status string
|
||||
AccountID string
|
||||
PlanType *string
|
||||
}
|
||||
|
||||
type Limit struct {
|
||||
LimitID string
|
||||
LimitName *string
|
||||
WindowType string
|
||||
UsedPercent float64
|
||||
WindowDurationMinutes int
|
||||
ResetsAt int64
|
||||
PlanType *string
|
||||
}
|
||||
|
||||
type UsageSummary struct {
|
||||
LifetimeTokens *int64
|
||||
PeakDailyTokens *int64
|
||||
LongestRunningTurnSec *int64
|
||||
CurrentStreakDays *int
|
||||
LongestStreakDays *int
|
||||
}
|
||||
|
||||
type UsagePoint struct {
|
||||
Date string
|
||||
TotalTokens int64
|
||||
}
|
||||
|
||||
type ResetCredits struct {
|
||||
AvailableCount int
|
||||
ExpiresAt []int64
|
||||
}
|
||||
|
||||
type Snapshot struct {
|
||||
Auth AuthFile
|
||||
Limits []Limit
|
||||
Summary UsageSummary
|
||||
Usage []UsagePoint
|
||||
ProfileAvailable bool
|
||||
UsageAvailable bool
|
||||
ResetCredits *ResetCredits
|
||||
FetchedAt time.Time
|
||||
}
|
||||
|
||||
func New(baseURL, managementKey string) *Client {
|
||||
return NewWithHTTPClient(baseURL, managementKey, &http.Client{Timeout: 15 * time.Second})
|
||||
}
|
||||
|
||||
func NewWithHTTPClient(baseURL, managementKey string, httpClient *http.Client) *Client {
|
||||
if httpClient == nil {
|
||||
httpClient = &http.Client{Timeout: 15 * time.Second}
|
||||
}
|
||||
return &Client{
|
||||
baseURL: strings.TrimRight(strings.TrimSpace(baseURL), "/"),
|
||||
managementKey: managementKey,
|
||||
httpClient: httpClient,
|
||||
now: time.Now,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) Configured() bool {
|
||||
return c.baseURL != "" && c.managementKey != ""
|
||||
}
|
||||
|
||||
func (c *Client) Auth(ctx context.Context, authIndex string) (AuthFile, error) {
|
||||
if !c.Configured() {
|
||||
return AuthFile{}, ErrNotConfigured
|
||||
}
|
||||
authIndex = strings.TrimSpace(authIndex)
|
||||
if authIndex == "" {
|
||||
return AuthFile{}, errors.New("authIndex 不能为空")
|
||||
}
|
||||
u, err := url.Parse(c.baseURL + "/v0/management/auth-files")
|
||||
if err != nil {
|
||||
return AuthFile{}, errors.New("CLIProxyAPI 地址无效")
|
||||
}
|
||||
q := u.Query()
|
||||
q.Set("auth_index", authIndex)
|
||||
u.RawQuery = q.Encode()
|
||||
body, err := c.managementRequest(ctx, http.MethodGet, u.String(), nil)
|
||||
if err != nil {
|
||||
return AuthFile{}, err
|
||||
}
|
||||
var raw any
|
||||
if err := json.Unmarshal(body, &raw); err != nil {
|
||||
return AuthFile{}, errors.New("CLIProxyAPI auth-files 响应格式错误")
|
||||
}
|
||||
items := authItems(raw)
|
||||
matches := make([]AuthFile, 0, 1)
|
||||
for _, item := range items {
|
||||
auth, ok := parseAuthFile(item, authIndex)
|
||||
if ok {
|
||||
matches = append(matches, auth)
|
||||
}
|
||||
}
|
||||
if len(matches) == 0 {
|
||||
return AuthFile{}, errors.New("未找到可用的 Codex auth")
|
||||
}
|
||||
if len(matches) != 1 {
|
||||
return AuthFile{}, errors.New("Codex auth 匹配结果不唯一")
|
||||
}
|
||||
return matches[0], nil
|
||||
}
|
||||
|
||||
func (c *Client) Snapshot(ctx context.Context, authIndex string) (Snapshot, error) {
|
||||
auth, err := c.Auth(ctx, authIndex)
|
||||
if err != nil {
|
||||
return Snapshot{}, err
|
||||
}
|
||||
receivedAt := c.now().UTC()
|
||||
usageBody, err := c.apiCall(ctx, auth, usageURL)
|
||||
if err != nil {
|
||||
return Snapshot{}, fmt.Errorf("获取用量限额失败: %w", err)
|
||||
}
|
||||
limits, availableFallback, effectivePlan, err := parseUsage(usageBody, auth.PlanType, receivedAt)
|
||||
if err != nil {
|
||||
return Snapshot{}, err
|
||||
}
|
||||
auth.PlanType = effectivePlan
|
||||
var summary UsageSummary
|
||||
usage := []UsagePoint{}
|
||||
profileAvailable := false
|
||||
usageAvailable := false
|
||||
if profileBody, profileErr := c.apiCall(ctx, auth, profileURL); profileErr == nil {
|
||||
if parsedSummary, parsedUsage, parsedUsageAvailable, parseErr := parseProfile(profileBody); parseErr == nil {
|
||||
summary = parsedSummary
|
||||
usage = parsedUsage
|
||||
usageAvailable = parsedUsageAvailable
|
||||
profileAvailable = usageAvailable || usageSummaryAvailable(summary)
|
||||
}
|
||||
}
|
||||
var resetCredits *ResetCredits
|
||||
resetCtx, cancelReset := context.WithTimeout(ctx, 5*time.Second)
|
||||
resetBody, resetErr := c.apiCall(resetCtx, auth, resetCreditsURL)
|
||||
cancelReset()
|
||||
if resetErr == nil {
|
||||
resetCredits = parseResetCredits(resetBody, availableFallback)
|
||||
} else if availableFallback > 0 {
|
||||
resetCredits = &ResetCredits{AvailableCount: availableFallback, ExpiresAt: []int64{}}
|
||||
}
|
||||
return Snapshot{
|
||||
Auth: auth,
|
||||
Limits: limits,
|
||||
Summary: summary,
|
||||
Usage: usage,
|
||||
ProfileAvailable: profileAvailable,
|
||||
UsageAvailable: usageAvailable,
|
||||
ResetCredits: resetCredits,
|
||||
FetchedAt: receivedAt,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *Client) managementRequest(ctx context.Context, method, endpoint string, body io.Reader) ([]byte, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, method, endpoint, body)
|
||||
if err != nil {
|
||||
return nil, errors.New("创建 CLIProxyAPI 请求失败")
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+c.managementKey)
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
httpClient := *c.httpClient
|
||||
httpClient.CheckRedirect = func(_ *http.Request, _ []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
}
|
||||
resp, err := httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, errors.New("CLIProxyAPI 请求失败")
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
payload, err := readLimited(resp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return nil, fmt.Errorf("CLIProxyAPI Management API 返回状态码 %d", resp.StatusCode)
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func (c *Client) apiCall(ctx context.Context, auth AuthFile, upstreamURL string) ([]byte, error) {
|
||||
requestBody := struct {
|
||||
AuthIndex string `json:"auth_index"`
|
||||
Method string `json:"method"`
|
||||
URL string `json:"url"`
|
||||
Header map[string]string `json:"header"`
|
||||
}{
|
||||
AuthIndex: auth.AuthIndex,
|
||||
Method: http.MethodGet,
|
||||
URL: upstreamURL,
|
||||
Header: map[string]string{
|
||||
"Authorization": "Bearer $TOKEN$",
|
||||
"Chatgpt-Account-Id": auth.AccountID,
|
||||
"Accept": "application/json",
|
||||
"Content-Type": "application/json",
|
||||
"User-Agent": codexUserAgent,
|
||||
},
|
||||
}
|
||||
encoded, err := json.Marshal(requestBody)
|
||||
if err != nil {
|
||||
return nil, errors.New("创建 CLIProxyAPI api-call 请求失败")
|
||||
}
|
||||
payload, err := c.managementRequest(ctx, http.MethodPost, c.baseURL+"/v0/management/api-call", strings.NewReader(string(encoded)))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var envelope struct {
|
||||
StatusCode int `json:"status_code"`
|
||||
Body json.RawMessage `json:"body"`
|
||||
}
|
||||
if err := json.Unmarshal(payload, &envelope); err != nil || envelope.StatusCode == 0 || len(envelope.Body) == 0 {
|
||||
return nil, errors.New("CLIProxyAPI api-call 响应格式错误")
|
||||
}
|
||||
if envelope.StatusCode < 200 || envelope.StatusCode >= 300 {
|
||||
if errorType := safeUpstreamErrorType(envelope.Body); errorType != "" {
|
||||
return nil, fmt.Errorf("CLIProxyAPI 上游请求返回状态码 %d (%s)", envelope.StatusCode, errorType)
|
||||
}
|
||||
return nil, fmt.Errorf("CLIProxyAPI 上游请求返回状态码 %d", envelope.StatusCode)
|
||||
}
|
||||
var bodyString string
|
||||
if len(envelope.Body) > 0 && envelope.Body[0] == '"' {
|
||||
if err := json.Unmarshal(envelope.Body, &bodyString); err != nil {
|
||||
return nil, errors.New("CLIProxyAPI api-call body 格式错误")
|
||||
}
|
||||
if len(bodyString) > maxResponseBody {
|
||||
return nil, errors.New("CLIProxyAPI 响应体过大")
|
||||
}
|
||||
return []byte(bodyString), nil
|
||||
}
|
||||
if len(envelope.Body) > maxResponseBody {
|
||||
return nil, errors.New("CLIProxyAPI 响应体过大")
|
||||
}
|
||||
return envelope.Body, nil
|
||||
}
|
||||
|
||||
func safeUpstreamErrorType(raw json.RawMessage) string {
|
||||
body := []byte(raw)
|
||||
if len(body) > 0 && body[0] == '"' {
|
||||
var text string
|
||||
if json.Unmarshal(body, &text) != nil {
|
||||
return ""
|
||||
}
|
||||
body = []byte(text)
|
||||
}
|
||||
var payload map[string]any
|
||||
if json.Unmarshal(body, &payload) != nil {
|
||||
return ""
|
||||
}
|
||||
candidates := []string{stringValue(payload, "type", "code")}
|
||||
if nested, ok := objectValue(payload, "error"); ok {
|
||||
candidates = append([]string{stringValue(nested, "type", "code")}, candidates...)
|
||||
}
|
||||
for _, candidate := range candidates {
|
||||
candidate = strings.TrimSpace(candidate)
|
||||
if candidate == "" || len(candidate) > 80 {
|
||||
continue
|
||||
}
|
||||
safe := true
|
||||
for _, r := range candidate {
|
||||
if !(r >= 'a' && r <= 'z' || r >= 'A' && r <= 'Z' || r >= '0' && r <= '9' || r == '_' || r == '-' || r == '.') {
|
||||
safe = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if safe {
|
||||
return candidate
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func readLimited(r io.Reader) ([]byte, error) {
|
||||
payload, err := io.ReadAll(io.LimitReader(r, maxResponseBody+1))
|
||||
if err != nil {
|
||||
return nil, errors.New("读取 CLIProxyAPI 响应失败")
|
||||
}
|
||||
if len(payload) > maxResponseBody {
|
||||
return nil, errors.New("CLIProxyAPI 响应体过大")
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func authItems(raw any) []map[string]any {
|
||||
switch value := raw.(type) {
|
||||
case []any:
|
||||
out := make([]map[string]any, 0, len(value))
|
||||
for _, item := range value {
|
||||
if object, ok := item.(map[string]any); ok {
|
||||
out = append(out, object)
|
||||
}
|
||||
}
|
||||
return out
|
||||
case map[string]any:
|
||||
for _, key := range []string{"auth_files", "authFiles", "files", "data", "items"} {
|
||||
if nested, ok := value[key]; ok {
|
||||
if out := authItems(nested); len(out) > 0 {
|
||||
return out
|
||||
}
|
||||
}
|
||||
}
|
||||
if stringValue(value, "auth_index", "authIndex") != "" {
|
||||
return []map[string]any{value}
|
||||
}
|
||||
}
|
||||
return []map[string]any{}
|
||||
}
|
||||
|
||||
func parseAuthFile(item map[string]any, expectedIndex string) (AuthFile, bool) {
|
||||
index := strings.TrimSpace(stringValue(item, "auth_index", "authIndex"))
|
||||
if index != expectedIndex {
|
||||
return AuthFile{}, false
|
||||
}
|
||||
provider := strings.ToLower(strings.TrimSpace(stringValue(item, "provider")))
|
||||
typ := strings.ToLower(strings.TrimSpace(stringValue(item, "type")))
|
||||
if provider != "" && provider != "codex" || typ != "" && typ != "codex" || provider == "" && typ == "" {
|
||||
return AuthFile{}, false
|
||||
}
|
||||
status := strings.TrimSpace(stringValue(item, "status"))
|
||||
statusLower := strings.ToLower(status)
|
||||
// Do not reject CPA's transient unavailable/error state: quota exhaustion
|
||||
// itself can set it, and this dashboard still needs to read the reset time.
|
||||
if boolValue(item, "disabled") || statusLower == "disabled" {
|
||||
return AuthFile{}, false
|
||||
}
|
||||
claims := claimsFrom(item)
|
||||
authClaims, _ := objectValue(claims, "https://api.openai.com/auth")
|
||||
accountID := strings.TrimSpace(stringValue(claims, "chatgpt_account_id", "chatgptAccountId"))
|
||||
if accountID == "" {
|
||||
accountID = strings.TrimSpace(stringValue(authClaims, "chatgpt_account_id", "chatgptAccountId"))
|
||||
}
|
||||
if accountID == "" {
|
||||
return AuthFile{}, false
|
||||
}
|
||||
var email *string
|
||||
if value := strings.TrimSpace(stringValue(item, "email")); value != "" {
|
||||
email = &value
|
||||
} else if value := strings.TrimSpace(stringValue(claims, "email")); value != "" {
|
||||
email = &value
|
||||
}
|
||||
var planType *string
|
||||
value := strings.TrimSpace(stringValue(claims, "plan_type", "planType"))
|
||||
if value == "" {
|
||||
value = strings.TrimSpace(stringValue(authClaims, "chatgpt_plan_type", "plan_type", "planType"))
|
||||
}
|
||||
if value != "" {
|
||||
planType = &value
|
||||
}
|
||||
return AuthFile{
|
||||
AuthIndex: index,
|
||||
Label: strings.TrimSpace(stringValue(item, "label")),
|
||||
Name: strings.TrimSpace(stringValue(item, "name")),
|
||||
Email: email,
|
||||
Status: status,
|
||||
AccountID: accountID,
|
||||
PlanType: planType,
|
||||
}, true
|
||||
}
|
||||
|
||||
func claimsFrom(item map[string]any) map[string]any {
|
||||
for _, key := range []string{"id_token_claims", "idTokenClaims", "id_token", "idToken"} {
|
||||
value, ok := item[key]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if object, ok := value.(map[string]any); ok {
|
||||
if nested, ok := object["claims"].(map[string]any); ok {
|
||||
return nested
|
||||
}
|
||||
return object
|
||||
}
|
||||
if encoded, ok := value.(string); ok {
|
||||
if claims := decodeClaimsString(encoded); claims != nil {
|
||||
return claims
|
||||
}
|
||||
}
|
||||
}
|
||||
return map[string]any{}
|
||||
}
|
||||
|
||||
func decodeClaimsString(value string) map[string]any {
|
||||
var claims map[string]any
|
||||
if json.Unmarshal([]byte(value), &claims) == nil {
|
||||
return claims
|
||||
}
|
||||
parts := strings.Split(value, ".")
|
||||
if len(parts) < 2 {
|
||||
return nil
|
||||
}
|
||||
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||
if err != nil || json.Unmarshal(payload, &claims) != nil {
|
||||
return nil
|
||||
}
|
||||
return claims
|
||||
}
|
||||
|
||||
func parseUsage(body []byte, fallbackPlan *string, receivedAt time.Time) ([]Limit, int, *string, error) {
|
||||
var root map[string]any
|
||||
if err := json.Unmarshal(body, &root); err != nil {
|
||||
return nil, 0, nil, errors.New("wham usage 响应格式错误")
|
||||
}
|
||||
if data, ok := objectValue(root, "data"); ok {
|
||||
root = data
|
||||
}
|
||||
planType := fallbackPlan
|
||||
if value := strings.TrimSpace(stringValue(root, "plan_type", "planType")); value != "" {
|
||||
planType = &value
|
||||
}
|
||||
limits := make([]Limit, 0)
|
||||
if raw, ok := objectValue(root, "rate_limit", "rateLimit"); ok {
|
||||
limits = append(limits, parseLimit("codex", nil, raw, planType, receivedAt)...)
|
||||
}
|
||||
if raw, ok := objectValue(root, "code_review_rate_limit", "codeReviewRateLimit"); ok {
|
||||
limits = append(limits, parseLimit("code_review", nil, raw, planType, receivedAt)...)
|
||||
}
|
||||
for _, raw := range arrayValue(root, "additional_rate_limits", "additionalRateLimits") {
|
||||
object, ok := raw.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
limitID := strings.TrimSpace(stringValue(object, "metered_feature", "meteredFeature"))
|
||||
if limitID == "" {
|
||||
continue
|
||||
}
|
||||
var name *string
|
||||
if value := strings.TrimSpace(stringValue(object, "limit_name", "limitName")); value != "" {
|
||||
name = &value
|
||||
}
|
||||
rateLimit := object
|
||||
if nested, ok := objectValue(object, "rate_limit", "rateLimit"); ok {
|
||||
rateLimit = nested
|
||||
}
|
||||
limits = append(limits, parseLimit(limitID, name, rateLimit, planType, receivedAt)...)
|
||||
}
|
||||
availableCount := 0
|
||||
if credits, ok := objectValue(root, "rate_limit_reset_credits", "rateLimitResetCredits"); ok {
|
||||
availableCount = int(numberValue(credits, "available_count", "availableCount"))
|
||||
}
|
||||
return limits, availableCount, planType, nil
|
||||
}
|
||||
|
||||
func parseLimit(limitID string, name *string, raw map[string]any, planType *string, receivedAt time.Time) []Limit {
|
||||
out := make([]Limit, 0, 2)
|
||||
for _, window := range []struct {
|
||||
kind string
|
||||
keys []string
|
||||
}{{"primary", []string{"primary_window", "primaryWindow"}}, {"secondary", []string{"secondary_window", "secondaryWindow"}}} {
|
||||
object, ok := objectValue(raw, window.keys...)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
used, usedOK := numericValue(object, "used_percent", "usedPercent")
|
||||
if !usedOK {
|
||||
continue
|
||||
}
|
||||
seconds := int64(numberValue(object, "limit_window_seconds", "limitWindowSeconds"))
|
||||
resetsAt := unixTimeValue(object, "reset_at", "resetAt")
|
||||
if resetsAt == 0 {
|
||||
resetAfter := int64(numberValue(object, "reset_after_seconds", "resetAfterSeconds"))
|
||||
if resetAfter > 0 {
|
||||
resetsAt = receivedAt.Unix() + resetAfter
|
||||
}
|
||||
}
|
||||
used = math.Max(0, math.Min(100, used))
|
||||
out = append(out, Limit{
|
||||
LimitID: limitID,
|
||||
LimitName: name,
|
||||
WindowType: window.kind,
|
||||
UsedPercent: used,
|
||||
WindowDurationMinutes: int((seconds + 59) / 60),
|
||||
ResetsAt: resetsAt,
|
||||
PlanType: planType,
|
||||
})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func parseProfile(body []byte) (UsageSummary, []UsagePoint, bool, error) {
|
||||
var root map[string]any
|
||||
if err := json.Unmarshal(body, &root); err != nil {
|
||||
return UsageSummary{}, nil, false, errors.New("wham profile 响应格式错误")
|
||||
}
|
||||
stats := root
|
||||
nestedStats := false
|
||||
if nested, ok := objectValue(root, "stats", "usage_stats", "usageStats"); ok {
|
||||
stats = nested
|
||||
nestedStats = true
|
||||
}
|
||||
summary := UsageSummary{
|
||||
LifetimeTokens: nonNegativeInt64Pointer(stats, "lifetime_tokens", "lifetimeTokens"),
|
||||
PeakDailyTokens: nonNegativeInt64Pointer(stats, "peak_daily_tokens", "peakDailyTokens"),
|
||||
LongestRunningTurnSec: nonNegativeInt64Pointer(stats, "longest_running_turn_sec", "longestRunningTurnSec"),
|
||||
CurrentStreakDays: nonNegativeIntPointer(stats, "current_streak_days", "currentStreakDays"),
|
||||
LongestStreakDays: nonNegativeIntPointer(stats, "longest_streak_days", "longestStreakDays"),
|
||||
}
|
||||
buckets, usageAvailable := arrayField(root, "daily_usage_buckets", "dailyUsageBuckets")
|
||||
if !usageAvailable && nestedStats {
|
||||
buckets, usageAvailable = arrayField(stats, "daily_usage_buckets", "dailyUsageBuckets")
|
||||
}
|
||||
usage := make([]UsagePoint, 0, len(buckets))
|
||||
for _, raw := range buckets {
|
||||
bucket, ok := raw.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
date := strings.TrimSpace(stringValue(bucket, "date", "start_date", "startDate"))
|
||||
if len(date) >= 10 {
|
||||
date = date[:10]
|
||||
}
|
||||
if _, err := time.Parse("2006-01-02", date); err != nil {
|
||||
continue
|
||||
}
|
||||
tokens, ok := numericValue(bucket, "total_tokens", "totalTokens", "tokens")
|
||||
if !ok || tokens < 0 {
|
||||
continue
|
||||
}
|
||||
usage = append(usage, UsagePoint{Date: date, TotalTokens: int64(tokens)})
|
||||
}
|
||||
return summary, usage, usageAvailable, nil
|
||||
}
|
||||
|
||||
func usageSummaryAvailable(summary UsageSummary) bool {
|
||||
return summary.LifetimeTokens != nil || summary.PeakDailyTokens != nil || summary.LongestRunningTurnSec != nil || summary.CurrentStreakDays != nil || summary.LongestStreakDays != nil
|
||||
}
|
||||
|
||||
func parseResetCredits(body []byte, fallback int) *ResetCredits {
|
||||
var root map[string]any
|
||||
if json.Unmarshal(body, &root) != nil {
|
||||
if fallback > 0 {
|
||||
return &ResetCredits{AvailableCount: fallback, ExpiresAt: []int64{}}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if data, ok := objectValue(root, "data"); ok {
|
||||
root = data
|
||||
}
|
||||
countPresent := hasValue(root, "available_count", "availableCount")
|
||||
count := int(numberValue(root, "available_count", "availableCount"))
|
||||
if countPresent && count <= 0 {
|
||||
return nil
|
||||
}
|
||||
if !countPresent {
|
||||
count = fallback
|
||||
}
|
||||
if count <= 0 {
|
||||
return nil
|
||||
}
|
||||
expiresAt := make([]int64, 0)
|
||||
for _, raw := range arrayValue(root, "credits", "available_credits", "availableCredits") {
|
||||
credit, ok := raw.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if status := strings.ToLower(strings.TrimSpace(stringValue(credit, "status"))); status != "" && status != "available" {
|
||||
continue
|
||||
}
|
||||
if value := unixTimeValue(credit, "expires_at", "expiresAt"); value > 0 {
|
||||
expiresAt = append(expiresAt, value)
|
||||
}
|
||||
}
|
||||
return &ResetCredits{AvailableCount: count, ExpiresAt: uniqueSorted(expiresAt)}
|
||||
}
|
||||
|
||||
func uniqueSorted(values []int64) []int64 {
|
||||
for i := 0; i < len(values); i++ {
|
||||
for j := i + 1; j < len(values); j++ {
|
||||
if values[j] < values[i] {
|
||||
values[i], values[j] = values[j], values[i]
|
||||
}
|
||||
}
|
||||
}
|
||||
out := values[:0]
|
||||
for _, value := range values {
|
||||
if len(out) == 0 || out[len(out)-1] != value {
|
||||
out = append(out, value)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func objectValue(object map[string]any, keys ...string) (map[string]any, bool) {
|
||||
for _, key := range keys {
|
||||
if value, ok := object[key].(map[string]any); ok {
|
||||
return value, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func arrayValue(object map[string]any, keys ...string) []any {
|
||||
value, _ := arrayField(object, keys...)
|
||||
return value
|
||||
}
|
||||
|
||||
func arrayField(object map[string]any, keys ...string) ([]any, bool) {
|
||||
for _, key := range keys {
|
||||
if value, ok := object[key].([]any); ok {
|
||||
return value, true
|
||||
}
|
||||
}
|
||||
return []any{}, false
|
||||
}
|
||||
|
||||
func stringValue(object map[string]any, keys ...string) string {
|
||||
for _, key := range keys {
|
||||
value, ok := object[key]
|
||||
if !ok || value == nil {
|
||||
continue
|
||||
}
|
||||
switch typed := value.(type) {
|
||||
case string:
|
||||
return typed
|
||||
case json.Number:
|
||||
return typed.String()
|
||||
case float64:
|
||||
if typed == math.Trunc(typed) {
|
||||
return strconv.FormatInt(int64(typed), 10)
|
||||
}
|
||||
return strconv.FormatFloat(typed, 'f', -1, 64)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func hasValue(object map[string]any, keys ...string) bool {
|
||||
for _, key := range keys {
|
||||
if value, ok := object[key]; ok && value != nil {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func boolValue(object map[string]any, keys ...string) bool {
|
||||
for _, key := range keys {
|
||||
if value, ok := object[key].(bool); ok {
|
||||
return value
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func numberValue(object map[string]any, keys ...string) float64 {
|
||||
value, _ := numericValue(object, keys...)
|
||||
return value
|
||||
}
|
||||
|
||||
func numericValue(object map[string]any, keys ...string) (float64, bool) {
|
||||
for _, key := range keys {
|
||||
value, ok := object[key]
|
||||
if !ok || value == nil {
|
||||
continue
|
||||
}
|
||||
var parsed float64
|
||||
var err error
|
||||
switch typed := value.(type) {
|
||||
case float64:
|
||||
parsed = typed
|
||||
case string:
|
||||
parsed, err = strconv.ParseFloat(strings.TrimSpace(typed), 64)
|
||||
case json.Number:
|
||||
parsed, err = typed.Float64()
|
||||
default:
|
||||
continue
|
||||
}
|
||||
if err == nil && !math.IsNaN(parsed) && !math.IsInf(parsed, 0) {
|
||||
return parsed, true
|
||||
}
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func nonNegativeInt64Pointer(object map[string]any, keys ...string) *int64 {
|
||||
if value, ok := numericValue(object, keys...); ok && value >= 0 {
|
||||
parsed := int64(value)
|
||||
return &parsed
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func nonNegativeIntPointer(object map[string]any, keys ...string) *int {
|
||||
value := nonNegativeInt64Pointer(object, keys...)
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
parsed := int(*value)
|
||||
return &parsed
|
||||
}
|
||||
|
||||
func unixTimeValue(object map[string]any, keys ...string) int64 {
|
||||
for _, key := range keys {
|
||||
value, ok := object[key]
|
||||
if !ok || value == nil {
|
||||
continue
|
||||
}
|
||||
if text, ok := value.(string); ok {
|
||||
if parsed, err := time.Parse(time.RFC3339, text); err == nil {
|
||||
return parsed.Unix()
|
||||
}
|
||||
if parsed, err := strconv.ParseInt(text, 10, 64); err == nil {
|
||||
return parsed
|
||||
}
|
||||
}
|
||||
return int64(numberValue(object, key))
|
||||
}
|
||||
return 0
|
||||
}
|
||||
@@ -0,0 +1,319 @@
|
||||
package cliproxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func authResponse(items string) string {
|
||||
return `{"auth_files":` + items + `}`
|
||||
}
|
||||
|
||||
func validAuth(index string) string {
|
||||
return `{"auth_index":"` + index + `","provider":"codex","type":"codex","label":"Main","email":"user@example.com","status":"ready","id_token":{"chatgpt_account_id":"acct-123","plan_type":"new_unknown_plan"}}`
|
||||
}
|
||||
|
||||
func TestAuthFiltersProviderStatusAndRequiresUniqueMatch(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
want string
|
||||
}{
|
||||
{name: "valid", body: authResponse(`[` + validAuth("auth/one") + `]`), want: ""},
|
||||
{name: "wrong provider", body: authResponse(`[{"auth_index":"auth/one","provider":"gemini","type":"gemini","id_token":{"chatgpt_account_id":"acct"}}]`), want: "未找到"},
|
||||
{name: "disabled", body: authResponse(`[{"auth_index":"auth/one","provider":"codex","disabled":true,"id_token":{"chatgpt_account_id":"acct"}}]`), want: "未找到"},
|
||||
{name: "quota unavailable remains readable", body: authResponse(`[{"auth_index":"auth/one","type":"codex","status":"error","unavailable":true,"id_token":{"chatgpt_account_id":"acct-123","plan_type":"new_unknown_plan"}}]`), want: ""},
|
||||
{name: "duplicate", body: authResponse(`[` + validAuth("auth/one") + `,` + validAuth("auth/one") + `]`), want: "不唯一"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if got := r.URL.Query().Get("auth_index"); got != "auth/one" {
|
||||
t.Fatalf("auth_index = %q", got)
|
||||
}
|
||||
_, _ = w.Write([]byte(tt.body))
|
||||
}))
|
||||
defer server.Close()
|
||||
client := New(server.URL, "management-secret")
|
||||
auth, err := client.Auth(context.Background(), "auth/one")
|
||||
if tt.want == "" {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if auth.AuthIndex != "auth/one" || auth.AccountID != "acct-123" || auth.PlanType == nil || *auth.PlanType != "new_unknown_plan" {
|
||||
t.Fatalf("auth = %#v", auth)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err == nil || !strings.Contains(err.Error(), tt.want) {
|
||||
t.Fatalf("error = %v; want %q", err, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSnapshotUsesManagementAuthorizationAndAPICallRequestStructure(t *testing.T) {
|
||||
var mu sync.Mutex
|
||||
var upstreamURLs []string
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if got := r.Header.Get("Authorization"); got != "Bearer management-secret" {
|
||||
t.Fatalf("Authorization = %q", got)
|
||||
}
|
||||
switch r.URL.Path {
|
||||
case "/v0/management/auth-files":
|
||||
_, _ = w.Write([]byte(authResponse(`[` + validAuth("auth 1") + `]`)))
|
||||
case "/v0/management/api-call":
|
||||
var request struct {
|
||||
AuthIndex string `json:"auth_index"`
|
||||
Method string `json:"method"`
|
||||
URL string `json:"url"`
|
||||
Header map[string]string `json:"header"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&request); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if request.AuthIndex != "auth 1" || request.Method != http.MethodGet || request.Header["Authorization"] != "Bearer $TOKEN$" || request.Header["Chatgpt-Account-Id"] != "acct-123" || request.Header["Accept"] != "application/json" || request.Header["Content-Type"] != "application/json" || request.Header["User-Agent"] != codexUserAgent {
|
||||
t.Fatalf("api-call request = %#v", request)
|
||||
}
|
||||
mu.Lock()
|
||||
upstreamURLs = append(upstreamURLs, request.URL)
|
||||
mu.Unlock()
|
||||
var body string
|
||||
switch request.URL {
|
||||
case usageURL:
|
||||
body = `{"rate_limit":{"primary_window":{"used_percent":12.75,"limit_window_seconds":301,"reset_after_seconds":60}}}`
|
||||
case profileURL:
|
||||
body = `{"lifetime_tokens":1000,"peak_daily_tokens":250,"longest_running_turn_sec":90,"current_streak_days":3,"longest_streak_days":8,"daily_usage_buckets":[{"start_date":"2026-08-14","tokens":77}]}`
|
||||
case resetCreditsURL:
|
||||
body = `{"available_count":1,"credits":[{"expires_at":"2026-08-20T00:00:00Z"}]}`
|
||||
default:
|
||||
t.Fatalf("unexpected upstream URL %q", request.URL)
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 200, "body": body})
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := New(server.URL, "management-secret")
|
||||
fixedNow := time.Date(2026, 8, 14, 12, 0, 0, 0, time.UTC)
|
||||
client.now = func() time.Time { return fixedNow }
|
||||
snapshot, err := client.Snapshot(context.Background(), "auth 1")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(snapshot.Limits) != 1 || snapshot.Limits[0].UsedPercent != 12.75 || snapshot.Limits[0].WindowDurationMinutes != 6 || snapshot.Limits[0].ResetsAt != fixedNow.Unix()+60 {
|
||||
t.Fatalf("limits = %#v", snapshot.Limits)
|
||||
}
|
||||
if snapshot.Limits[0].PlanType == nil || *snapshot.Limits[0].PlanType != "new_unknown_plan" {
|
||||
t.Fatalf("plan type was not preserved: %#v", snapshot.Limits[0].PlanType)
|
||||
}
|
||||
if snapshot.Summary.LifetimeTokens == nil || *snapshot.Summary.LifetimeTokens != 1000 || snapshot.Summary.PeakDailyTokens == nil || *snapshot.Summary.PeakDailyTokens != 250 || len(snapshot.Usage) != 1 || snapshot.Usage[0].TotalTokens != 77 {
|
||||
t.Fatalf("profile = %#v usage = %#v", snapshot.Summary, snapshot.Usage)
|
||||
}
|
||||
if snapshot.ResetCredits == nil || snapshot.ResetCredits.AvailableCount != 1 || len(snapshot.ResetCredits.ExpiresAt) != 1 || snapshot.ResetCredits.ExpiresAt[0] != time.Date(2026, 8, 20, 0, 0, 0, 0, time.UTC).Unix() {
|
||||
t.Fatalf("reset credits = %#v", snapshot.ResetCredits)
|
||||
}
|
||||
if len(upstreamURLs) != 3 {
|
||||
t.Fatalf("upstream URLs = %#v", upstreamURLs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUsageParsesAllLimitKindsAndClampsPercentages(t *testing.T) {
|
||||
plan := "enterprise-new"
|
||||
body := []byte(`{
|
||||
"rate_limit":{"primary_window":{"used_percent":-1.25,"limit_window_seconds":300,"reset_at":1800000000}},
|
||||
"code_review_rate_limit":{"secondary_window":{"used_percent":101.5,"limit_window_seconds":604801,"reset_at":1800000100}},
|
||||
"additional_rate_limits":[{"metered_feature":"spark","limit_name":"Spark usage","primary_window":{"used_percent":45.5,"limit_window_seconds":61,"reset_at":1800000200}}],
|
||||
"rate_limit_reset_credits":{"available_count":2}
|
||||
}`)
|
||||
limits, fallback, effectivePlan, err := parseUsage(body, &plan, time.Now())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(limits) != 3 || limits[0].LimitID != "codex" || limits[0].UsedPercent != 0 || limits[1].LimitID != "code_review" || limits[1].UsedPercent != 100 || limits[1].WindowDurationMinutes != 10081 || limits[2].LimitID != "spark" || limits[2].LimitName == nil || *limits[2].LimitName != "Spark usage" || limits[2].WindowDurationMinutes != 2 || fallback != 2 || effectivePlan == nil || *effectivePlan != plan {
|
||||
t.Fatalf("limits = %#v fallback = %d", limits, fallback)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUsageSkipsWindowsWithoutUsedPercentage(t *testing.T) {
|
||||
limits, _, _, err := parseUsage([]byte(`{"rate_limit":{"primary_window":{"limit_window_seconds":18000,"reset_at":1800000000}}}`), nil, time.Now())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(limits) != 0 {
|
||||
t.Fatalf("limits = %#v; missing used_percent must not become zero usage", limits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUsageReadsNestedAdditionalLimitsAndPlanFromPayload(t *testing.T) {
|
||||
body := []byte(`{"data":{"plan_type":"prolite","additional_rate_limits":[{"metered_feature":"codex_bengalfox","limit_name":"GPT-5.3-Codex-Spark","rate_limit":{"primary_window":{"used_percent":9.5,"limit_window_seconds":18000,"reset_at":1800000000}}}]}}`)
|
||||
limits, _, plan, err := parseUsage(body, nil, time.Now())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(limits) != 1 || limits[0].LimitID != "codex_bengalfox" || limits[0].UsedPercent != 9.5 || plan == nil || *plan != "prolite" || limits[0].PlanType == nil || *limits[0].PlanType != "prolite" {
|
||||
t.Fatalf("limits = %#v plan = %v", limits, plan)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileMissingOrNullBucketsAreNotMarkedAvailable(t *testing.T) {
|
||||
for _, body := range []string{`{}`, `{"stats":{}}`, `{"stats":{"daily_usage_buckets":null}}`} {
|
||||
summary, usage, usageAvailable, err := parseProfile([]byte(body))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if usageAvailable || usage == nil || len(usage) != 0 || usageSummaryAvailable(summary) {
|
||||
t.Fatalf("body=%s summary=%#v usage=%#v available=%v", body, summary, usage, usageAvailable)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileInvalidOptionalMetricsRemainUnavailable(t *testing.T) {
|
||||
summary, usage, usageAvailable, err := parseProfile([]byte(`{"stats":{"lifetime_tokens":"unknown","daily_usage_buckets":[{"start_date":"2026-08-14","tokens":"bad"}]}}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if summary.LifetimeTokens != nil || !usageAvailable || len(usage) != 0 {
|
||||
t.Fatalf("summary = %#v usage = %#v", summary, usage)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResetCreditsIgnoreNonAvailableDetails(t *testing.T) {
|
||||
credits := parseResetCredits([]byte(`{"available_count":2,"credits":[{"status":"redeemed","expires_at":"2026-08-19T00:00:00Z"},{"status":"available","expires_at":"2026-08-20T00:00:00Z"}]}`), 0)
|
||||
if credits == nil || credits.AvailableCount != 2 || len(credits.ExpiresAt) != 1 || credits.ExpiresAt[0] != time.Date(2026, 8, 20, 0, 0, 0, 0, time.UTC).Unix() {
|
||||
t.Fatalf("credits = %#v", credits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileFailureDoesNotHideRateLimits(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/v0/management/auth-files" {
|
||||
_, _ = w.Write([]byte(authResponse(`[` + validAuth("auth-1") + `]`)))
|
||||
return
|
||||
}
|
||||
var request struct {
|
||||
URL string `json:"url"`
|
||||
}
|
||||
_ = json.NewDecoder(r.Body).Decode(&request)
|
||||
switch request.URL {
|
||||
case usageURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 200, "body": `{"rate_limit":{"primary_window":{"used_percent":50,"limit_window_seconds":18000,"reset_at":1800000000}}}`})
|
||||
case profileURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 503, "body": `{"error":"profile unavailable"}`})
|
||||
case resetCreditsURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 404, "body": `{}`})
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
snapshot, err := New(server.URL, "management-secret").Snapshot(context.Background(), "auth-1")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(snapshot.Limits) != 1 || snapshot.ProfileAvailable || snapshot.Usage == nil || len(snapshot.Usage) != 0 || snapshot.Summary.LifetimeTokens != nil {
|
||||
t.Fatalf("snapshot = %#v", snapshot)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResetCreditsSuccessfulZeroOverridesUsageFallback(t *testing.T) {
|
||||
if credits := parseResetCredits([]byte(`{"available_count":0,"credits":[]}`), 3); credits != nil {
|
||||
t.Fatalf("credits = %#v; successful detail response is authoritative", credits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOptionalResetFailureFallsBackToUsageCount(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/v0/management/auth-files" {
|
||||
_, _ = w.Write([]byte(authResponse(`[` + validAuth("auth-1") + `]`)))
|
||||
return
|
||||
}
|
||||
var request struct {
|
||||
URL string `json:"url"`
|
||||
}
|
||||
_ = json.NewDecoder(r.Body).Decode(&request)
|
||||
switch request.URL {
|
||||
case usageURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 200, "body": `{"rate_limit_reset_credits":{"available_count":3}}`})
|
||||
case profileURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 200, "body": `{}`})
|
||||
case resetCreditsURL:
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 503, "body": `{"secret":"upstream detail"}`})
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
client := New(server.URL, "management-secret")
|
||||
snapshot, err := client.Snapshot(context.Background(), "auth-1")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if snapshot.ResetCredits == nil || snapshot.ResetCredits.AvailableCount != 3 || snapshot.ResetCredits.ExpiresAt == nil {
|
||||
t.Fatalf("reset credits = %#v", snapshot.ResetCredits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagementRequestDoesNotFollowRedirectsWithSecret(t *testing.T) {
|
||||
reachedRedirect := false
|
||||
target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
reachedRedirect = true
|
||||
if r.Header.Get("Authorization") != "" {
|
||||
t.Fatal("management Authorization header reached redirect target")
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer target.Close()
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, target.URL, http.StatusFound)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := New(server.URL, "management-secret")
|
||||
_, err := client.Auth(context.Background(), "auth-1")
|
||||
if err == nil || !strings.Contains(err.Error(), "302") {
|
||||
t.Fatalf("error = %v", err)
|
||||
}
|
||||
if reachedRedirect {
|
||||
t.Fatal("redirect was unexpectedly followed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeUpstreamErrorTypeOnlyReturnsBoundedIdentifiers(t *testing.T) {
|
||||
if got := safeUpstreamErrorType(json.RawMessage(`"{\"error\":{\"type\":\"token_expired\"}}"`)); got != "token_expired" {
|
||||
t.Fatalf("type = %q", got)
|
||||
}
|
||||
if got := safeUpstreamErrorType(json.RawMessage(`{"error":{"type":"secret bearer value"}}`)); got != "" {
|
||||
t.Fatalf("unsafe type = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpstreamErrorDoesNotLeakBodyOrManagementKey(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/v0/management/auth-files" {
|
||||
_, _ = w.Write([]byte(authResponse(`[` + validAuth("auth-1") + `]`)))
|
||||
return
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"status_code": 401, "body": `{"access_token":"raw-token","detail":"private detail"}`})
|
||||
}))
|
||||
defer server.Close()
|
||||
client := New(server.URL, "management-secret")
|
||||
_, err := client.Snapshot(context.Background(), "auth-1")
|
||||
if err == nil {
|
||||
t.Fatal("snapshot unexpectedly succeeded")
|
||||
}
|
||||
message := err.Error()
|
||||
for _, secret := range []string{"raw-token", "private detail", "management-secret"} {
|
||||
if strings.Contains(message, secret) {
|
||||
t.Fatalf("error leaked %q: %s", secret, message)
|
||||
}
|
||||
}
|
||||
if !strings.Contains(message, "401") {
|
||||
t.Fatalf("error = %q", message)
|
||||
}
|
||||
}
|
||||
@@ -1,175 +0,0 @@
|
||||
package codex
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Client struct {
|
||||
mu sync.Mutex
|
||||
cmd *exec.Cmd
|
||||
in io.WriteCloser
|
||||
pending map[int64]chan envelope
|
||||
id atomic.Int64
|
||||
connected bool
|
||||
configDir string
|
||||
notify func(string, json.RawMessage)
|
||||
}
|
||||
type envelope struct {
|
||||
ID *int64 `json:"id,omitempty"`
|
||||
Method string `json:"method,omitempty"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
Result json.RawMessage `json:"result,omitempty"`
|
||||
Error any `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
func New(configDir string, notify func(string, json.RawMessage)) *Client {
|
||||
return &Client{pending: map[int64]chan envelope{}, configDir: configDir, notify: notify}
|
||||
}
|
||||
func (c *Client) Start(ctx context.Context) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.connected {
|
||||
return nil
|
||||
}
|
||||
if e := ensureConfigDir(c.configDir); e != nil {
|
||||
return e
|
||||
}
|
||||
cmd := exec.CommandContext(ctx, "codex", "app-server")
|
||||
cmd.Env = append(os.Environ(), "CODEX_HOME="+c.configDir)
|
||||
out, e := cmd.StdoutPipe()
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
in, e := cmd.StdinPipe()
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
cmd.Stderr = os.Stderr
|
||||
if e = cmd.Start(); e != nil {
|
||||
return e
|
||||
}
|
||||
c.cmd, c.in, c.connected = cmd, in, true
|
||||
go c.read(cmd, out)
|
||||
go func() { _ = cmd.Wait(); c.failAll(cmd) }()
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensureConfigDir(dir string) error {
|
||||
if e := os.MkdirAll(dir, 0700); e != nil {
|
||||
return fmt.Errorf("create CODEX_HOME %q: %w", dir, e)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (c *Client) Initialize(ctx context.Context) error {
|
||||
var out any
|
||||
if e := c.Call(ctx, "initialize", map[string]any{"clientInfo": map[string]any{"name": "codex-helper", "title": "Codex Helper", "version": "0.1.0"}, "capabilities": map[string]any{}}, &out); e != nil {
|
||||
return e
|
||||
}
|
||||
return c.send(map[string]any{"method": "initialized", "params": map[string]any{}})
|
||||
}
|
||||
func (c *Client) read(cmd *exec.Cmd, r io.Reader) {
|
||||
s := bufio.NewScanner(r)
|
||||
s.Buffer(make([]byte, 64*1024), 8*1024*1024)
|
||||
for s.Scan() {
|
||||
var e envelope
|
||||
if json.Unmarshal(s.Bytes(), &e) != nil {
|
||||
continue
|
||||
}
|
||||
if e.ID != nil {
|
||||
c.mu.Lock()
|
||||
ch := c.pending[*e.ID]
|
||||
delete(c.pending, *e.ID)
|
||||
c.mu.Unlock()
|
||||
if ch != nil {
|
||||
ch <- e
|
||||
}
|
||||
} else if e.Method != "" && c.notify != nil {
|
||||
go c.notify(e.Method, e.Params)
|
||||
}
|
||||
}
|
||||
c.failAll(cmd)
|
||||
}
|
||||
func (c *Client) failAll(cmd *exec.Cmd) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
// A previous process may finish after its replacement has started. It must
|
||||
// not mark the new connection as disconnected or fail its pending calls.
|
||||
if c.cmd != cmd {
|
||||
return
|
||||
}
|
||||
c.connected = false
|
||||
c.cmd = nil
|
||||
c.in = nil
|
||||
for id, ch := range c.pending {
|
||||
ch <- envelope{Error: "app-server disconnected"}
|
||||
delete(c.pending, id)
|
||||
}
|
||||
}
|
||||
func (c *Client) send(v any) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if !c.connected {
|
||||
return errors.New("app-server unavailable")
|
||||
}
|
||||
b, _ := json.Marshal(v)
|
||||
b = append(b, '\n')
|
||||
_, e := c.in.Write(b)
|
||||
return e
|
||||
}
|
||||
func (c *Client) Call(ctx context.Context, method string, params any, out any) error {
|
||||
id := c.id.Add(1)
|
||||
ch := make(chan envelope, 1)
|
||||
c.mu.Lock()
|
||||
c.pending[id] = ch
|
||||
c.mu.Unlock()
|
||||
if e := c.send(map[string]any{"id": id, "method": method, "params": params}); e != nil {
|
||||
return e
|
||||
}
|
||||
select {
|
||||
case e := <-ch:
|
||||
if e.Error != nil {
|
||||
return fmt.Errorf("app-server %s: %v", method, e.Error)
|
||||
}
|
||||
if out != nil {
|
||||
return json.Unmarshal(e.Result, out)
|
||||
}
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
c.mu.Lock()
|
||||
delete(c.pending, id)
|
||||
c.mu.Unlock()
|
||||
return ctx.Err()
|
||||
case <-time.After(20 * time.Second):
|
||||
c.mu.Lock()
|
||||
delete(c.pending, id)
|
||||
c.mu.Unlock()
|
||||
return errors.New("app-server timeout")
|
||||
}
|
||||
}
|
||||
func (c *Client) Close() error {
|
||||
c.mu.Lock()
|
||||
cmd := c.cmd
|
||||
c.cmd = nil
|
||||
c.in = nil
|
||||
c.connected = false
|
||||
for id, ch := range c.pending {
|
||||
ch <- envelope{Error: "app-server disconnected"}
|
||||
delete(c.pending, id)
|
||||
}
|
||||
c.mu.Unlock()
|
||||
if cmd != nil && cmd.Process != nil {
|
||||
return cmd.Process.Kill()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (c *Client) Connected() bool { c.mu.Lock(); defer c.mu.Unlock(); return c.connected }
|
||||
@@ -1,41 +0,0 @@
|
||||
package codex
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEnsureConfigDirCreatesNestedDirectory(t *testing.T) {
|
||||
dir := filepath.Join(t.TempDir(), "accounts", "2", "codex")
|
||||
if err := ensureConfigDir(dir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
info, err := os.Stat(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !info.IsDir() {
|
||||
t.Fatalf("%s is not a directory", dir)
|
||||
}
|
||||
if got := info.Mode().Perm(); got != 0700 {
|
||||
t.Fatalf("permissions = %o; want 700", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureConfigDirReportsCreationFailure(t *testing.T) {
|
||||
parent := t.TempDir()
|
||||
file := filepath.Join(parent, "not-a-directory")
|
||||
if err := os.WriteFile(file, []byte("x"), 0600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dir := filepath.Join(file, "codex")
|
||||
err := ensureConfigDir(dir)
|
||||
if err == nil {
|
||||
t.Fatal("expected directory creation to fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "create CODEX_HOME") || !strings.Contains(err.Error(), dir) {
|
||||
t.Fatalf("error = %q; want CODEX_HOME path context", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestAuthIndexCRUDAndUniqueness(t *testing.T) {
|
||||
s, err := Open(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer s.DB.Close()
|
||||
first, err := s.CreateAccount("First", "auth-1", "personal")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if first.AuthIndex != "auth-1" {
|
||||
t.Fatalf("first = %#v", first)
|
||||
}
|
||||
if _, err = s.CreateAccount("Duplicate", "auth-1"); err == nil {
|
||||
t.Fatal("duplicate auth_index unexpectedly succeeded")
|
||||
}
|
||||
second, err := s.CreateAccountWithVisibility("Second", "auth-2", "team", true)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newIndex := "auth-3"
|
||||
visible := false
|
||||
if err = s.UpdateAccountBinding(second.ID, "Updated", &newIndex, "any", &visible); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
updated, err := s.Account(second.ID)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if updated.AuthIndex != "auth-3" || updated.DisplayName != "Updated" || updated.PublicVisible || updated.ExpectedKind != "any" {
|
||||
t.Fatalf("updated = %#v", updated)
|
||||
}
|
||||
conflict := "auth-1"
|
||||
if err = s.UpdateAccountBinding(second.ID, "Updated", &conflict, "any", nil); err == nil {
|
||||
t.Fatal("conflicting update unexpectedly succeeded")
|
||||
}
|
||||
stillUpdated, err := s.Account(second.ID)
|
||||
if err != nil || stillUpdated.AuthIndex != "auth-3" {
|
||||
t.Fatalf("failed update changed binding: %#v err=%v", stillUpdated, err)
|
||||
}
|
||||
used, err := s.AuthIndexUsed("auth-1", first.ID)
|
||||
if err != nil || used {
|
||||
t.Fatalf("exclude current account: used=%v err=%v", used, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthIndexMigrationIsIdempotentAndPreservesAccounts(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
db, err := sql.Open("sqlite", filepath.Join(dir, "codex-helper.db"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = db.Exec(`CREATE TABLE accounts (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT, display_name TEXT NOT NULL, email TEXT,
|
||||
plan_type TEXT, expected_kind TEXT NOT NULL DEFAULT 'any', public_visible INTEGER NOT NULL DEFAULT 0,
|
||||
connected INTEGER NOT NULL DEFAULT 0, created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL
|
||||
); INSERT INTO accounts VALUES(7,'Legacy','legacy@example.com','plus','personal',1,1,1,2);`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = db.Close()
|
||||
|
||||
for i := 0; i < 2; i++ {
|
||||
s, openErr := Open(dir)
|
||||
if openErr != nil {
|
||||
t.Fatal(openErr)
|
||||
}
|
||||
account, accountErr := s.Account(7)
|
||||
if accountErr != nil {
|
||||
_ = s.DB.Close()
|
||||
t.Fatal(accountErr)
|
||||
}
|
||||
if account.AuthIndex != "" || account.DisplayName != "Legacy" || account.Email == nil || *account.Email != "legacy@example.com" || !account.PublicVisible {
|
||||
_ = s.DB.Close()
|
||||
t.Fatalf("migrated account = %#v", account)
|
||||
}
|
||||
var indexCount int
|
||||
if queryErr := s.DB.QueryRow("SELECT COUNT(*) FROM sqlite_master WHERE type='index' AND name='idx_accounts_auth_index'").Scan(&indexCount); queryErr != nil || indexCount != 1 {
|
||||
_ = s.DB.Close()
|
||||
t.Fatalf("index count = %d err=%v", indexCount, queryErr)
|
||||
}
|
||||
_ = s.DB.Close()
|
||||
}
|
||||
}
|
||||
@@ -94,6 +94,7 @@ func (s *Store) migrateAccounts() error {
|
||||
display_name TEXT NOT NULL,
|
||||
email TEXT,
|
||||
plan_type TEXT,
|
||||
auth_index TEXT NOT NULL DEFAULT '',
|
||||
expected_kind TEXT NOT NULL DEFAULT 'any',
|
||||
public_visible INTEGER NOT NULL DEFAULT 0,
|
||||
connected INTEGER NOT NULL DEFAULT 0,
|
||||
@@ -102,7 +103,7 @@ func (s *Store) migrateAccounts() error {
|
||||
)`); err != nil {
|
||||
return err
|
||||
}
|
||||
var hasExpectedKind, hasPublicVisible bool
|
||||
var hasAuthIndex, hasExpectedKind, hasPublicVisible bool
|
||||
rows, qerr := tx.Query("PRAGMA table_info(accounts)")
|
||||
if qerr != nil {
|
||||
return qerr
|
||||
@@ -112,10 +113,16 @@ func (s *Store) migrateAccounts() error {
|
||||
var name, typ string
|
||||
var def any
|
||||
_ = rows.Scan(&cid, &name, &typ, ¬null, &def, &pk)
|
||||
hasAuthIndex = hasAuthIndex || name == "auth_index"
|
||||
hasExpectedKind = hasExpectedKind || name == "expected_kind"
|
||||
hasPublicVisible = hasPublicVisible || name == "public_visible"
|
||||
}
|
||||
rows.Close()
|
||||
if !hasAuthIndex {
|
||||
if _, err = tx.Exec("ALTER TABLE accounts ADD COLUMN auth_index TEXT NOT NULL DEFAULT ''"); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if !hasExpectedKind {
|
||||
if _, err = tx.Exec("ALTER TABLE accounts ADD COLUMN expected_kind TEXT NOT NULL DEFAULT 'any'"); err != nil {
|
||||
return err
|
||||
@@ -126,13 +133,31 @@ func (s *Store) migrateAccounts() error {
|
||||
return err
|
||||
}
|
||||
}
|
||||
var count int
|
||||
if _, err = tx.Exec("CREATE UNIQUE INDEX IF NOT EXISTS idx_accounts_auth_index ON accounts(auth_index) WHERE auth_index <> ''"); err != nil {
|
||||
return err
|
||||
}
|
||||
// Legacy local app-server accounts have no CPA binding. Preserve their
|
||||
// metadata and history, but do not report them as connected until an
|
||||
// administrator assigns a valid authIndex.
|
||||
if _, err = tx.Exec("UPDATE accounts SET connected=0 WHERE auth_index=''"); err != nil {
|
||||
return err
|
||||
}
|
||||
var count, legacyRows int
|
||||
if err = tx.QueryRow("SELECT COUNT(*) FROM accounts").Scan(&count); err != nil {
|
||||
return err
|
||||
}
|
||||
if count == 0 {
|
||||
if _, err = tx.Exec("INSERT INTO accounts(id,display_name,created_at,updated_at) VALUES(1,'默认账号',?,?)", time.Now().Unix(), time.Now().Unix()); err != nil {
|
||||
return err
|
||||
for _, table := range []string{"daily_usage", "limit_snapshots"} {
|
||||
var tableRows int
|
||||
if err = tx.QueryRow("SELECT COUNT(*) FROM " + table).Scan(&tableRows); err != nil {
|
||||
return err
|
||||
}
|
||||
legacyRows += tableRows
|
||||
}
|
||||
if legacyRows > 0 {
|
||||
if _, err = tx.Exec("INSERT INTO accounts(id,display_name,created_at,updated_at) VALUES(1,'默认账号',?,?)", time.Now().Unix(), time.Now().Unix()); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, table := range []string{"daily_usage", "limit_snapshots"} {
|
||||
@@ -187,6 +212,7 @@ type Account struct {
|
||||
DisplayName string `json:"displayName"`
|
||||
Email *string `json:"email"`
|
||||
PlanType *string `json:"planType"`
|
||||
AuthIndex string `json:"authIndex,omitempty"`
|
||||
ExpectedKind string `json:"expectedKind"`
|
||||
PublicVisible bool `json:"publicVisible"`
|
||||
ActualKind string `json:"actualKind"`
|
||||
@@ -198,7 +224,7 @@ type Account struct {
|
||||
}
|
||||
|
||||
func (s *Store) Accounts() ([]Account, error) {
|
||||
rows, e := s.DB.Query("SELECT id,display_name,email,plan_type,expected_kind,public_visible,connected,created_at,updated_at FROM accounts ORDER BY id")
|
||||
rows, e := s.DB.Query("SELECT id,display_name,email,plan_type,auth_index,expected_kind,public_visible,connected,created_at,updated_at FROM accounts ORDER BY id")
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
@@ -206,7 +232,7 @@ func (s *Store) Accounts() ([]Account, error) {
|
||||
out := []Account{}
|
||||
for rows.Next() {
|
||||
var a Account
|
||||
if e = rows.Scan(&a.ID, &a.DisplayName, &a.Email, &a.PlanType, &a.ExpectedKind, &a.PublicVisible, &a.Connected, &a.CreatedAt, &a.UpdatedAt); e != nil {
|
||||
if e = rows.Scan(&a.ID, &a.DisplayName, &a.Email, &a.PlanType, &a.AuthIndex, &a.ExpectedKind, &a.PublicVisible, &a.Connected, &a.CreatedAt, &a.UpdatedAt); e != nil {
|
||||
return nil, e
|
||||
}
|
||||
a.ActualKind, a.ValidationStatus = AccountKind(a.PlanType), validationStatus(a.ExpectedKind, a.Connected, a.PlanType)
|
||||
@@ -228,8 +254,8 @@ func (s *Store) Accounts() ([]Account, error) {
|
||||
|
||||
func (s *Store) Account(id int64) (Account, error) {
|
||||
var a Account
|
||||
err := s.DB.QueryRow("SELECT id,display_name,email,plan_type,expected_kind,public_visible,connected,created_at,updated_at FROM accounts WHERE id=?", id).
|
||||
Scan(&a.ID, &a.DisplayName, &a.Email, &a.PlanType, &a.ExpectedKind, &a.PublicVisible, &a.Connected, &a.CreatedAt, &a.UpdatedAt)
|
||||
err := s.DB.QueryRow("SELECT id,display_name,email,plan_type,auth_index,expected_kind,public_visible,connected,created_at,updated_at FROM accounts WHERE id=?", id).
|
||||
Scan(&a.ID, &a.DisplayName, &a.Email, &a.PlanType, &a.AuthIndex, &a.ExpectedKind, &a.PublicVisible, &a.Connected, &a.CreatedAt, &a.UpdatedAt)
|
||||
if err != nil {
|
||||
return Account{}, err
|
||||
}
|
||||
@@ -237,38 +263,46 @@ func (s *Store) Account(id int64) (Account, error) {
|
||||
return a, nil
|
||||
}
|
||||
|
||||
func (s *Store) CreateAccount(name string, kinds ...string) (Account, error) {
|
||||
return s.createAccount(name, false, kinds...)
|
||||
func (s *Store) CreateAccount(name, authIndex string, kinds ...string) (Account, error) {
|
||||
return s.createAccount(name, authIndex, false, kinds...)
|
||||
}
|
||||
|
||||
func (s *Store) CreateAccountWithVisibility(name, expectedKind string, publicVisible bool) (Account, error) {
|
||||
return s.createAccount(name, publicVisible, expectedKind)
|
||||
func (s *Store) CreateAccountWithVisibility(name, authIndex, expectedKind string, publicVisible bool) (Account, error) {
|
||||
return s.createAccount(name, authIndex, publicVisible, expectedKind)
|
||||
}
|
||||
|
||||
func (s *Store) createAccount(name string, publicVisible bool, kinds ...string) (Account, error) {
|
||||
func (s *Store) createAccount(name, authIndex string, publicVisible bool, kinds ...string) (Account, error) {
|
||||
expectedKind := "any"
|
||||
if len(kinds) > 0 {
|
||||
expectedKind = kinds[0]
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
r, e := s.DB.Exec("INSERT INTO accounts(display_name,expected_kind,public_visible,created_at,updated_at) VALUES(?,?,?,?,?)", name, expectedKind, publicVisible, now, now)
|
||||
r, e := s.DB.Exec("INSERT INTO accounts(display_name,auth_index,expected_kind,public_visible,created_at,updated_at) VALUES(?,?,?,?,?,?)", name, authIndex, expectedKind, publicVisible, now, now)
|
||||
if e != nil {
|
||||
return Account{}, e
|
||||
}
|
||||
id, _ := r.LastInsertId()
|
||||
return Account{ID: id, DisplayName: name, ExpectedKind: expectedKind, PublicVisible: publicVisible, ActualKind: "unknown", ValidationStatus: "pending", CreatedAt: now, UpdatedAt: now}, nil
|
||||
return Account{ID: id, DisplayName: name, AuthIndex: authIndex, ExpectedKind: expectedKind, PublicVisible: publicVisible, ActualKind: "unknown", ValidationStatus: "pending", CreatedAt: now, UpdatedAt: now}, nil
|
||||
}
|
||||
func (s *Store) UpdateAccountSettings(id int64, name, expectedKind string) error {
|
||||
return s.UpdateAccountSettingsWithVisibility(id, name, expectedKind, nil)
|
||||
}
|
||||
|
||||
func (s *Store) UpdateAccountSettingsWithVisibility(id int64, name, expectedKind string, publicVisible *bool) error {
|
||||
return s.UpdateAccountBinding(id, name, nil, expectedKind, publicVisible)
|
||||
}
|
||||
|
||||
func (s *Store) UpdateAccountBinding(id int64, name string, authIndex *string, expectedKind string, publicVisible *bool) error {
|
||||
var r sql.Result
|
||||
var e error
|
||||
if publicVisible == nil {
|
||||
if authIndex == nil && publicVisible == nil {
|
||||
r, e = s.DB.Exec("UPDATE accounts SET display_name=?,expected_kind=?,updated_at=? WHERE id=?", name, expectedKind, time.Now().Unix(), id)
|
||||
} else {
|
||||
} else if authIndex == nil {
|
||||
r, e = s.DB.Exec("UPDATE accounts SET display_name=?,expected_kind=?,public_visible=?,updated_at=? WHERE id=?", name, expectedKind, *publicVisible, time.Now().Unix(), id)
|
||||
} else if publicVisible == nil {
|
||||
r, e = s.DB.Exec("UPDATE accounts SET display_name=?,auth_index=?,expected_kind=?,updated_at=? WHERE id=?", name, *authIndex, expectedKind, time.Now().Unix(), id)
|
||||
} else {
|
||||
r, e = s.DB.Exec("UPDATE accounts SET display_name=?,auth_index=?,expected_kind=?,public_visible=?,updated_at=? WHERE id=?", name, *authIndex, expectedKind, *publicVisible, time.Now().Unix(), id)
|
||||
}
|
||||
if e != nil {
|
||||
return e
|
||||
@@ -279,6 +313,12 @@ func (s *Store) UpdateAccountSettingsWithVisibility(id int64, name, expectedKind
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Store) AuthIndexUsed(authIndex string, excludeID int64) (bool, error) {
|
||||
var count int
|
||||
err := s.DB.QueryRow("SELECT COUNT(*) FROM accounts WHERE auth_index=? AND id<>?", authIndex, excludeID).Scan(&count)
|
||||
return count > 0, err
|
||||
}
|
||||
|
||||
func (s *Store) RenameAccount(id int64, name string) error {
|
||||
var kind string
|
||||
if err := s.DB.QueryRow("SELECT expected_kind FROM accounts WHERE id=?", id).Scan(&kind); err != nil {
|
||||
@@ -323,8 +363,18 @@ func (s *Store) UpdateAccount(id int64, email, plan *string, connected bool) err
|
||||
return e
|
||||
}
|
||||
func (s *Store) DeleteAccount(id int64) error {
|
||||
_, e := s.DB.Exec("DELETE FROM accounts WHERE id=?", id)
|
||||
return e
|
||||
tx, err := s.DB.Begin()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
if _, err = tx.Exec("DELETE FROM notifications WHERE dedupe_key GLOB ?", fmt.Sprintf("%d:*", id)); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err = tx.Exec("DELETE FROM accounts WHERE id=?", id); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
func (s *Store) Get(key string) (string, bool) {
|
||||
|
||||
@@ -3,6 +3,7 @@ package store
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
@@ -33,18 +34,26 @@ func TestAccountsAndPerAccountUsage(t *testing.T) {
|
||||
}
|
||||
defer s.DB.Close()
|
||||
accounts, err := s.Accounts()
|
||||
if err != nil || len(accounts) != 1 || accounts[0].ID != 1 || accounts[0].PublicVisible {
|
||||
t.Fatalf("default accounts = %#v, %v", accounts, err)
|
||||
if err != nil || len(accounts) != 0 {
|
||||
t.Fatalf("fresh accounts = %#v, %v; want no automatic default account", accounts, err)
|
||||
}
|
||||
second, err := s.CreateAccount("Team workspace")
|
||||
first, err := s.CreateAccount("Personal workspace", "auth-personal")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, id := range []int64{1, second.ID} {
|
||||
second, err := s.CreateAccount("Team workspace", "auth-team")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, id := range []int64{first.ID, second.ID} {
|
||||
if _, err = s.DB.Exec("INSERT INTO daily_usage(account_id,date,total_tokens,fetched_at) VALUES(?,?,?,?)", id, "2026-08-13", id*100, 1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if _, err = s.DB.Exec(`INSERT INTO notifications
|
||||
(dedupe_key,channel,kind,status,scheduled_at,body) VALUES(?, 'configured', 'after', 'pending', 1, '{}')`, fmt.Sprintf("%d:codex:primary:1:after", second.ID)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = s.DeleteAccount(second.ID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -52,6 +61,9 @@ func TestAccountsAndPerAccountUsage(t *testing.T) {
|
||||
if err = s.DB.QueryRow("SELECT COUNT(*) FROM daily_usage WHERE account_id=?", second.ID).Scan(&count); err != nil || count != 0 {
|
||||
t.Fatalf("usage was not cascaded: %d, %v", count, err)
|
||||
}
|
||||
if err = s.DB.QueryRow("SELECT COUNT(*) FROM notifications WHERE dedupe_key GLOB ?", fmt.Sprintf("%d:*", second.ID)).Scan(&count); err != nil || count != 0 {
|
||||
t.Fatalf("notifications were not deleted: %d, %v", count, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAccountKindAndValidation(t *testing.T) {
|
||||
@@ -97,7 +109,7 @@ func TestExistingAccountsGainExpectedKind(t *testing.T) {
|
||||
if err != nil || len(accounts) != 1 {
|
||||
t.Fatalf("accounts = %#v, %v", accounts, err)
|
||||
}
|
||||
if accounts[0].ExpectedKind != "any" || accounts[0].PublicVisible || accounts[0].ActualKind != "team" || accounts[0].ValidationStatus != "matched" {
|
||||
if accounts[0].AuthIndex != "" || accounts[0].ExpectedKind != "any" || accounts[0].PublicVisible || accounts[0].Connected || accounts[0].ActualKind != "team" || accounts[0].ValidationStatus != "pending" {
|
||||
t.Fatalf("migrated account = %#v", accounts[0])
|
||||
}
|
||||
}
|
||||
@@ -109,11 +121,11 @@ func TestAccountVisibilitySettings(t *testing.T) {
|
||||
}
|
||||
defer s.DB.Close()
|
||||
|
||||
private, err := s.CreateAccount("私有账号")
|
||||
private, err := s.CreateAccount("私有账号", "private-auth")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
public, err := s.CreateAccountWithVisibility("公开账号", "team", true)
|
||||
public, err := s.CreateAccountWithVisibility("公开账号", "public-auth", "team", true)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user