878 lines
29 KiB
Go
878 lines
29 KiB
Go
package app
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"os"
|
|
"path/filepath"
|
|
"sort"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"codex-helper/internal/cliproxy"
|
|
"codex-helper/internal/security"
|
|
"codex-helper/internal/store"
|
|
)
|
|
|
|
func (a *App) api(w http.ResponseWriter, r *http.Request) {
|
|
p := strings.TrimPrefix(r.URL.Path, "/api/v1/")
|
|
if p == "system/status" {
|
|
if r.Method != http.MethodGet {
|
|
jsonOut(w, http.StatusMethodNotAllowed, map[string]string{"error": "方法不允许"})
|
|
return
|
|
}
|
|
jsonOut(w, http.StatusOK, map[string]any{"initialized": a.store.Initialized(), "version": Version, "cpa": a.cpaConfigured()})
|
|
return
|
|
}
|
|
if p == "setup" && r.Method == http.MethodPost {
|
|
a.setup(w, r)
|
|
return
|
|
}
|
|
if p == "auth/login" && r.Method == http.MethodPost {
|
|
a.login(w, r)
|
|
return
|
|
}
|
|
if p == "accounts" && readOnlyMethod(r.Method) {
|
|
a.accountsAPI(w, r)
|
|
return
|
|
}
|
|
if p == "dashboard" && readOnlyMethod(r.Method) {
|
|
a.dashboardAPI(w, r)
|
|
return
|
|
}
|
|
if !a.require(w, r) {
|
|
return
|
|
}
|
|
switch {
|
|
case p == "auth/me":
|
|
var username string
|
|
_ = a.store.DB.QueryRow("SELECT username FROM admin WHERE id=1").Scan(&username)
|
|
jsonOut(w, http.StatusOK, map[string]string{"username": username})
|
|
case p == "auth/logout" && r.Method == http.MethodPost:
|
|
a.logout(w, r)
|
|
case p == "accounts" && r.Method == http.MethodPost:
|
|
a.createAccount(w, r)
|
|
case strings.HasPrefix(p, "accounts/"):
|
|
a.accountAPI(w, r, p)
|
|
case p == "dashboard":
|
|
a.dashboardAPI(w, r)
|
|
case p == "sync" && r.Method == http.MethodPost:
|
|
id, _ := strconv.ParseInt(r.URL.Query().Get("accountId"), 10, 64)
|
|
if id == 0 {
|
|
id = 1
|
|
}
|
|
if err := a.syncAccount(r.Context(), id); err != nil {
|
|
jsonOut(w, http.StatusBadGateway, map[string]string{"error": err.Error()})
|
|
} else {
|
|
jsonOut(w, http.StatusOK, map[string]bool{"ok": true})
|
|
}
|
|
case p == "settings/general":
|
|
a.generalAPI(w, r)
|
|
case p == "settings/smtp":
|
|
a.smtpAPI(w, r)
|
|
case p == "settings/smtp/test" && r.Method == http.MethodPost:
|
|
a.smtpTest(w, r)
|
|
case p == "settings/telegram":
|
|
a.telegramAPI(w, r)
|
|
case p == "settings/telegram/test" && r.Method == http.MethodPost:
|
|
a.telegramTest(w, r)
|
|
case p == "settings/telegram/bind" && r.Method == http.MethodPost:
|
|
a.telegramMu.Lock()
|
|
code := fmt.Sprintf("%06d", time.Now().UnixNano()%1000000)
|
|
_ = a.store.SetJSON("telegram_bind", map[string]any{"code": code, "expires": time.Now().Add(10 * time.Minute).Unix()})
|
|
a.telegramMu.Unlock()
|
|
jsonOut(w, http.StatusOK, map[string]string{"code": code})
|
|
case p == "maintenance/cleanup" && r.Method == http.MethodPost:
|
|
n, err := a.store.Cleanup(a.general().RetentionDays)
|
|
if err != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
} else {
|
|
jsonOut(w, http.StatusOK, map[string]int64{"deleted": n})
|
|
}
|
|
case p == "maintenance/backup":
|
|
dir, err := os.MkdirTemp(a.dataDir, "backup-")
|
|
if err != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
defer os.RemoveAll(dir)
|
|
path := filepath.Join(dir, "codex-helper.db")
|
|
if err = a.store.Backup(r.Context(), path); err != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
w.Header().Set("Content-Disposition", `attachment; filename="codex-helper.db"`)
|
|
http.ServeFile(w, r, path)
|
|
default:
|
|
jsonOut(w, http.StatusNotFound, map[string]string{"error": "接口不存在"})
|
|
}
|
|
}
|
|
|
|
func readOnlyMethod(method string) bool {
|
|
return method == http.MethodGet || method == http.MethodHead
|
|
}
|
|
|
|
func (a *App) accountsAPI(w http.ResponseWriter, r *http.Request) {
|
|
if !a.store.Initialized() {
|
|
jsonOut(w, http.StatusConflict, map[string]string{"error": "请先初始化"})
|
|
return
|
|
}
|
|
accounts, err := a.store.Accounts()
|
|
if err != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
if !a.authed(r) {
|
|
visible := make([]store.Account, 0, len(accounts))
|
|
for _, account := range accounts {
|
|
if account.PublicVisible {
|
|
visible = append(visible, publicAccount(account))
|
|
}
|
|
}
|
|
accounts = visible
|
|
}
|
|
jsonOut(w, http.StatusOK, accounts)
|
|
}
|
|
|
|
func publicAccount(account store.Account) store.Account {
|
|
account.Email = nil
|
|
account.AuthIndex = ""
|
|
account.ExpectedKind = "any"
|
|
account.ActualKind = "unknown"
|
|
account.ValidationStatus = "unknown"
|
|
account.PossibleDuplicate = false
|
|
account.CreatedAt = 0
|
|
account.UpdatedAt = 0
|
|
return account
|
|
}
|
|
|
|
func (a *App) dashboardAPI(w http.ResponseWriter, r *http.Request) {
|
|
if !a.store.Initialized() {
|
|
jsonOut(w, http.StatusConflict, map[string]string{"error": "请先初始化"})
|
|
return
|
|
}
|
|
id, _ := strconv.ParseInt(r.URL.Query().Get("accountId"), 10, 64)
|
|
if id == 0 {
|
|
id = 1
|
|
}
|
|
account, err := a.store.Account(id)
|
|
if err != nil {
|
|
if err == sql.ErrNoRows {
|
|
jsonOut(w, http.StatusNotFound, map[string]string{"error": "账号不存在"})
|
|
} else {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
}
|
|
return
|
|
}
|
|
if !a.authed(r) && !account.PublicVisible {
|
|
jsonOut(w, http.StatusNotFound, map[string]string{"error": "账号不存在"})
|
|
return
|
|
}
|
|
rt := a.runtime(id)
|
|
if rt == nil {
|
|
jsonOut(w, http.StatusNotFound, map[string]string{"error": "账号不存在"})
|
|
return
|
|
}
|
|
rt.syncing.Lock()
|
|
dashboard := rt.dash
|
|
rt.syncing.Unlock()
|
|
if dashboard.Limits == nil {
|
|
dashboard.Limits = []LimitBucket{}
|
|
}
|
|
if dashboard.Usage == nil {
|
|
dashboard.Usage = []UsagePoint{}
|
|
}
|
|
if !a.authed(r) {
|
|
dashboard = publicDashboard(dashboard)
|
|
}
|
|
jsonOut(w, http.StatusOK, dashboard)
|
|
}
|
|
|
|
func publicDashboard(d Dashboard) Dashboard {
|
|
d.Account.Email = nil
|
|
d.Account.AuthMode = nil
|
|
d.LastError = ""
|
|
return d
|
|
}
|
|
|
|
func (a *App) createAccount(w http.ResponseWriter, r *http.Request) {
|
|
var in struct {
|
|
DisplayName string `json:"displayName"`
|
|
AuthIndex string `json:"authIndex"`
|
|
ExpectedKind string `json:"expectedKind"`
|
|
PublicVisible bool `json:"publicVisible"`
|
|
}
|
|
if decode(r, &in) != nil {
|
|
jsonOut(w, http.StatusBadRequest, map[string]string{"error": "请求格式错误"})
|
|
return
|
|
}
|
|
in.AuthIndex = strings.TrimSpace(in.AuthIndex)
|
|
if in.AuthIndex == "" {
|
|
jsonOut(w, http.StatusBadRequest, map[string]string{"error": "authIndex 不能为空"})
|
|
return
|
|
}
|
|
if in.ExpectedKind == "" {
|
|
in.ExpectedKind = "any"
|
|
}
|
|
if !store.ValidExpectedKind(in.ExpectedKind) {
|
|
jsonOut(w, http.StatusBadRequest, map[string]string{"error": "连接类型无效"})
|
|
return
|
|
}
|
|
if used, err := a.store.AuthIndexUsed(in.AuthIndex, 0); err != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
} else if used {
|
|
jsonOut(w, http.StatusConflict, map[string]string{"error": "authIndex 已绑定"})
|
|
return
|
|
}
|
|
snapshot, err := a.fetchSnapshot(r.Context(), in.AuthIndex)
|
|
if err != nil {
|
|
jsonOut(w, http.StatusBadGateway, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
name := strings.TrimSpace(in.DisplayName)
|
|
if name == "" {
|
|
name = strings.TrimSpace(snapshot.Auth.Label)
|
|
}
|
|
if name == "" {
|
|
name = strings.TrimSpace(snapshot.Auth.Name)
|
|
}
|
|
if name == "" {
|
|
name = "新账号"
|
|
}
|
|
account, err := a.store.CreateAccountWithVisibility(name, in.AuthIndex, in.ExpectedKind, in.PublicVisible)
|
|
if err != nil {
|
|
if authIndexConflict(err) {
|
|
jsonOut(w, http.StatusConflict, map[string]string{"error": "authIndex 已绑定"})
|
|
} else {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
}
|
|
return
|
|
}
|
|
a.addRuntime(account.ID)
|
|
rt := a.runtime(account.ID)
|
|
rt.syncing.Lock()
|
|
dashboard := dashboardFromSnapshot(account, snapshot)
|
|
_, err = a.persistDashboard(dashboard)
|
|
if err == nil {
|
|
rt.dash = dashboard
|
|
}
|
|
rt.syncing.Unlock()
|
|
if err != nil {
|
|
a.mu.Lock()
|
|
delete(a.runtimes, account.ID)
|
|
a.mu.Unlock()
|
|
if rollbackErr := a.store.DeleteAccount(account.ID); rollbackErr != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": "账号创建失败且本地回滚失败"})
|
|
return
|
|
}
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
created, err := a.store.Account(account.ID)
|
|
if err != nil {
|
|
a.mu.Lock()
|
|
delete(a.runtimes, account.ID)
|
|
a.mu.Unlock()
|
|
if rollbackErr := a.store.DeleteAccount(account.ID); rollbackErr != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": "账号创建失败且本地回滚失败"})
|
|
return
|
|
}
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
jsonOut(w, http.StatusCreated, created)
|
|
}
|
|
|
|
func (a *App) accountAPI(w http.ResponseWriter, r *http.Request, p string) {
|
|
parts := strings.Split(p, "/")
|
|
if len(parts) < 2 {
|
|
jsonOut(w, http.StatusNotFound, map[string]string{"error": "接口不存在"})
|
|
return
|
|
}
|
|
id, err := strconv.ParseInt(parts[1], 10, 64)
|
|
if err != nil {
|
|
jsonOut(w, http.StatusBadRequest, map[string]string{"error": "账号 ID 无效"})
|
|
return
|
|
}
|
|
rt := a.runtime(id)
|
|
if rt == nil {
|
|
jsonOut(w, http.StatusNotFound, map[string]string{"error": "账号不存在"})
|
|
return
|
|
}
|
|
action := ""
|
|
if len(parts) > 2 {
|
|
action = parts[2]
|
|
}
|
|
switch {
|
|
case action == "" && r.Method == http.MethodPut:
|
|
a.updateAccount(w, r, id, rt)
|
|
case action == "" && r.Method == http.MethodDelete:
|
|
// Match the reminder -> global-runtime -> account-runtime lock order
|
|
// so no queued reminder for the deleted binding can be sent later.
|
|
a.reminderMu.Lock()
|
|
defer a.reminderMu.Unlock()
|
|
a.mu.Lock()
|
|
rt.syncing.Lock()
|
|
err := a.store.DeleteAccount(id)
|
|
if err == nil {
|
|
delete(a.runtimes, id)
|
|
}
|
|
rt.syncing.Unlock()
|
|
a.mu.Unlock()
|
|
if err != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
jsonOut(w, http.StatusOK, map[string]bool{"ok": true})
|
|
case action == "sync" && r.Method == http.MethodPost:
|
|
if err := a.syncAccount(r.Context(), id); err != nil {
|
|
jsonOut(w, http.StatusBadGateway, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
jsonOut(w, http.StatusOK, map[string]bool{"ok": true})
|
|
default:
|
|
jsonOut(w, http.StatusNotFound, map[string]string{"error": "接口不存在"})
|
|
}
|
|
}
|
|
|
|
func (a *App) updateAccount(w http.ResponseWriter, r *http.Request, id int64, rt *accountRuntime) {
|
|
var in struct {
|
|
DisplayName string `json:"displayName"`
|
|
AuthIndex *string `json:"authIndex"`
|
|
ExpectedKind string `json:"expectedKind"`
|
|
PublicVisible *bool `json:"publicVisible"`
|
|
}
|
|
if decode(r, &in) != nil || strings.TrimSpace(in.DisplayName) == "" {
|
|
jsonOut(w, http.StatusBadRequest, map[string]string{"error": "名称不能为空"})
|
|
return
|
|
}
|
|
account, err := a.store.Account(id)
|
|
if err != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
if in.ExpectedKind == "" {
|
|
in.ExpectedKind = account.ExpectedKind
|
|
}
|
|
if !store.ValidExpectedKind(in.ExpectedKind) {
|
|
jsonOut(w, http.StatusBadRequest, map[string]string{"error": "连接类型无效"})
|
|
return
|
|
}
|
|
name := strings.TrimSpace(in.DisplayName)
|
|
if in.AuthIndex == nil || strings.TrimSpace(*in.AuthIndex) == account.AuthIndex {
|
|
if err = a.store.UpdateAccountSettingsWithVisibility(id, name, in.ExpectedKind, in.PublicVisible); err != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
rt.syncing.Lock()
|
|
rt.dash.DisplayName = name
|
|
rt.syncing.Unlock()
|
|
jsonOut(w, http.StatusOK, map[string]bool{"ok": true})
|
|
return
|
|
}
|
|
newAuthIndex := strings.TrimSpace(*in.AuthIndex)
|
|
if newAuthIndex == "" {
|
|
jsonOut(w, http.StatusBadRequest, map[string]string{"error": "authIndex 不能为空"})
|
|
return
|
|
}
|
|
if used, checkErr := a.store.AuthIndexUsed(newAuthIndex, id); checkErr != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": checkErr.Error()})
|
|
return
|
|
} else if used {
|
|
jsonOut(w, http.StatusConflict, map[string]string{"error": "authIndex 已绑定"})
|
|
return
|
|
}
|
|
snapshot, err := a.fetchSnapshot(r.Context(), newAuthIndex)
|
|
if err != nil {
|
|
jsonOut(w, http.StatusBadGateway, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
updatedAccount := account
|
|
updatedAccount.DisplayName = name
|
|
updatedAccount.AuthIndex = newAuthIndex
|
|
updatedAccount.ExpectedKind = in.ExpectedKind
|
|
if in.PublicVisible != nil {
|
|
updatedAccount.PublicVisible = *in.PublicVisible
|
|
}
|
|
dashboard := dashboardFromSnapshot(updatedAccount, snapshot)
|
|
// Finish any old-identity reminder work before replacing the binding and
|
|
// deleting its dedupe records.
|
|
a.reminderMu.Lock()
|
|
defer a.reminderMu.Unlock()
|
|
rt.syncing.Lock()
|
|
defer rt.syncing.Unlock()
|
|
tx, err := a.store.DB.BeginTx(r.Context(), nil)
|
|
if err == nil {
|
|
var result sql.Result
|
|
if in.PublicVisible == nil {
|
|
result, err = tx.Exec("UPDATE accounts SET display_name=?,auth_index=?,expected_kind=?,updated_at=? WHERE id=?", name, newAuthIndex, in.ExpectedKind, time.Now().Unix(), id)
|
|
} else {
|
|
result, err = tx.Exec("UPDATE accounts SET display_name=?,auth_index=?,expected_kind=?,public_visible=?,updated_at=? WHERE id=?", name, newAuthIndex, in.ExpectedKind, *in.PublicVisible, time.Now().Unix(), id)
|
|
}
|
|
if err == nil {
|
|
var affected int64
|
|
affected, err = result.RowsAffected()
|
|
if err == nil && affected == 0 {
|
|
err = sql.ErrNoRows
|
|
}
|
|
}
|
|
}
|
|
resetDetected := false
|
|
if err == nil {
|
|
_, err = tx.Exec("DELETE FROM daily_usage WHERE account_id=?", id)
|
|
}
|
|
if err == nil {
|
|
_, err = tx.Exec("DELETE FROM limit_snapshots WHERE account_id=?", id)
|
|
}
|
|
if err == nil {
|
|
_, err = tx.Exec("DELETE FROM notifications WHERE dedupe_key GLOB ?", strconv.FormatInt(id, 10)+":*")
|
|
}
|
|
if err == nil {
|
|
resetDetected, err = a.persistDashboardTx(tx, dashboard)
|
|
}
|
|
if err == nil {
|
|
err = tx.Commit()
|
|
} else if tx != nil {
|
|
_ = tx.Rollback()
|
|
}
|
|
if err != nil {
|
|
if authIndexConflict(err) {
|
|
jsonOut(w, http.StatusConflict, map[string]string{"error": "authIndex 已绑定"})
|
|
} else {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
}
|
|
return
|
|
}
|
|
rt.dash = dashboard
|
|
if resetDetected {
|
|
go a.processReminders()
|
|
}
|
|
jsonOut(w, http.StatusOK, map[string]bool{"ok": true})
|
|
}
|
|
|
|
func authIndexConflict(err error) bool {
|
|
if err == nil {
|
|
return false
|
|
}
|
|
message := strings.ToLower(err.Error())
|
|
return strings.Contains(message, "unique") && strings.Contains(message, "accounts.auth_index")
|
|
}
|
|
|
|
func (a *App) setup(w http.ResponseWriter, r *http.Request) {
|
|
if a.store.Initialized() {
|
|
jsonOut(w, http.StatusConflict, map[string]string{"error": "系统已初始化"})
|
|
return
|
|
}
|
|
var in struct{ Username, Password, Timezone string }
|
|
if decode(r, &in) != nil || len(in.Username) < 3 || len(in.Password) < 10 {
|
|
jsonOut(w, http.StatusBadRequest, map[string]string{"error": "用户名至少3位,密码至少10位"})
|
|
return
|
|
}
|
|
tx, err := a.store.DB.Begin()
|
|
if err != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
defer tx.Rollback()
|
|
_, err = tx.Exec("INSERT INTO admin(id,username,password_hash,created_at) VALUES(1,?,?,?)", in.Username, security.Password(in.Password), time.Now().Unix())
|
|
if err == nil {
|
|
g := defaults()
|
|
if in.Timezone != "" {
|
|
if _, zoneErr := time.LoadLocation(in.Timezone); zoneErr == nil {
|
|
g.Timezone = in.Timezone
|
|
}
|
|
}
|
|
b, _ := json.Marshal(g)
|
|
_, err = tx.Exec("INSERT INTO settings(key,value,updated_at) VALUES('general',?,?),('initialized','true',?)", string(b), time.Now().Unix(), time.Now().Unix())
|
|
}
|
|
if err == nil {
|
|
err = tx.Commit()
|
|
}
|
|
if err != nil {
|
|
jsonOut(w, http.StatusInternalServerError, map[string]string{"error": err.Error()})
|
|
return
|
|
}
|
|
a.newSession(w, in.Username)
|
|
jsonOut(w, http.StatusCreated, map[string]bool{"ok": true})
|
|
}
|
|
|
|
func (a *App) login(w http.ResponseWriter, r *http.Request) {
|
|
if !a.store.Initialized() {
|
|
jsonOut(w, http.StatusConflict, map[string]string{"error": "请先初始化"})
|
|
return
|
|
}
|
|
ip := r.RemoteAddr
|
|
value, _ := a.loginAttempts.LoadOrStore(ip, []time.Time{})
|
|
attempts := value.([]time.Time)
|
|
now := time.Now()
|
|
fresh := attempts[:0]
|
|
for _, attempt := range attempts {
|
|
if now.Sub(attempt) < 15*time.Minute {
|
|
fresh = append(fresh, attempt)
|
|
}
|
|
}
|
|
if len(fresh) >= 10 {
|
|
jsonOut(w, http.StatusTooManyRequests, map[string]string{"error": "尝试次数过多,请稍后再试"})
|
|
return
|
|
}
|
|
var in struct{ Username, Password string }
|
|
_ = decode(r, &in)
|
|
var user, hash string
|
|
err := a.store.DB.QueryRow("SELECT username,password_hash FROM admin WHERE id=1").Scan(&user, &hash)
|
|
if err != nil || user != in.Username || !security.VerifyPassword(hash, in.Password) {
|
|
a.loginAttempts.Store(ip, append(fresh, now))
|
|
jsonOut(w, http.StatusUnauthorized, map[string]string{"error": "用户名或密码错误"})
|
|
return
|
|
}
|
|
a.loginAttempts.Delete(ip)
|
|
a.newSession(w, user)
|
|
jsonOut(w, http.StatusOK, map[string]bool{"ok": true})
|
|
}
|
|
|
|
func (a *App) newSession(w http.ResponseWriter, _ string) {
|
|
token := security.Random(32)
|
|
_, _ = a.store.DB.Exec("DELETE FROM sessions WHERE expires_at<?", time.Now().Unix())
|
|
_, _ = a.store.DB.Exec("INSERT INTO sessions(token_hash,expires_at,created_at) VALUES(?,?,?)", security.HashToken(token), time.Now().Add(7*24*time.Hour).Unix(), time.Now().Unix())
|
|
http.SetCookie(w, &http.Cookie{Name: "session", Value: token, Path: "/", HttpOnly: true, SameSite: http.SameSiteStrictMode, Secure: false, MaxAge: 604800})
|
|
}
|
|
|
|
func (a *App) logout(w http.ResponseWriter, r *http.Request) {
|
|
if cookie, err := r.Cookie("session"); err == nil {
|
|
_, _ = a.store.DB.Exec("DELETE FROM sessions WHERE token_hash=?", security.HashToken(cookie.Value))
|
|
}
|
|
http.SetCookie(w, &http.Cookie{Name: "session", Path: "/", MaxAge: -1, HttpOnly: true})
|
|
jsonOut(w, http.StatusOK, map[string]bool{"ok": true})
|
|
}
|
|
|
|
func (a *App) general() GeneralSettings {
|
|
g := defaults()
|
|
a.store.GetJSON("general", &g)
|
|
g.SyncMinutes = automaticSyncMinutes
|
|
if g.RetentionDays < 1 {
|
|
g.RetentionDays = 90
|
|
}
|
|
return g
|
|
}
|
|
|
|
func (a *App) generalAPI(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method == http.MethodGet {
|
|
jsonOut(w, http.StatusOK, a.general())
|
|
return
|
|
}
|
|
var g GeneralSettings
|
|
if decode(r, &g) != nil || g.SyncMinutes < 1 || g.SyncMinutes > 60 || g.RetentionDays < 30 || g.RetentionDays > 365 || g.BeforeMinutes < 1 || g.BeforeMinutes > 1440 {
|
|
jsonOut(w, http.StatusBadRequest, map[string]string{"error": "设置值不合法"})
|
|
return
|
|
}
|
|
if _, err := time.LoadLocation(g.Timezone); err != nil {
|
|
jsonOut(w, http.StatusBadRequest, map[string]string{"error": "无效时区"})
|
|
return
|
|
}
|
|
g.SyncMinutes = automaticSyncMinutes
|
|
_ = a.store.SetJSON("general", g)
|
|
jsonOut(w, http.StatusOK, g)
|
|
}
|
|
|
|
func (a *App) fetchSnapshot(ctx context.Context, authIndex string) (cliproxy.Snapshot, error) {
|
|
if !a.cpaConfigured() {
|
|
return cliproxy.Snapshot{}, cliproxy.ErrNotConfigured
|
|
}
|
|
return a.cpa.Snapshot(ctx, authIndex)
|
|
}
|
|
|
|
func (a *App) syncAccount(ctx context.Context, id int64) error {
|
|
rt := a.runtime(id)
|
|
if rt == nil {
|
|
return errorsForSync("账号不存在")
|
|
}
|
|
rt.syncing.Lock()
|
|
defer rt.syncing.Unlock()
|
|
account, err := a.store.Account(id)
|
|
if err != nil {
|
|
return a.markSyncFailure(rt, id, "账号不存在")
|
|
}
|
|
if strings.TrimSpace(account.AuthIndex) == "" {
|
|
return a.markSyncFailure(rt, id, "账号尚未绑定 authIndex")
|
|
}
|
|
snapshot, err := a.fetchSnapshot(ctx, account.AuthIndex)
|
|
if err != nil {
|
|
return a.markSyncFailure(rt, id, err.Error())
|
|
}
|
|
dashboard := dashboardFromSnapshot(account, snapshot)
|
|
mergeProfileData(&dashboard, rt.dash, snapshot.ProfileAvailable, snapshot.UsageAvailable)
|
|
resetDetected, err := a.persistDashboard(dashboard)
|
|
if err != nil {
|
|
return a.markSyncFailure(rt, id, err.Error())
|
|
}
|
|
rt.dash = dashboard
|
|
if resetDetected {
|
|
go a.processReminders()
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func errorsForSync(message string) error {
|
|
return fmt.Errorf("%s", message)
|
|
}
|
|
|
|
func (a *App) markSyncFailure(rt *accountRuntime, id int64, message string) error {
|
|
if rt.dash.AccountID == 0 {
|
|
rt.dash.AccountID = id
|
|
}
|
|
if rt.dash.Limits == nil {
|
|
rt.dash.Limits = []LimitBucket{}
|
|
}
|
|
if rt.dash.Usage == nil {
|
|
rt.dash.Usage = []UsagePoint{}
|
|
}
|
|
rt.dash.Stale = true
|
|
rt.dash.LastError = message
|
|
return errorsForSync(message)
|
|
}
|
|
|
|
func mergeProfileData(current *Dashboard, previous Dashboard, profileAvailable, usageAvailable bool) {
|
|
if !profileAvailable {
|
|
current.Summary = previous.Summary
|
|
} else {
|
|
if current.Summary.LifetimeTokens == nil {
|
|
current.Summary.LifetimeTokens = previous.Summary.LifetimeTokens
|
|
}
|
|
if current.Summary.LongestRunningTurnSec == nil {
|
|
current.Summary.LongestRunningTurnSec = previous.Summary.LongestRunningTurnSec
|
|
}
|
|
if current.Summary.CurrentStreakDays == nil {
|
|
current.Summary.CurrentStreakDays = previous.Summary.CurrentStreakDays
|
|
}
|
|
if current.Summary.LongestStreakDays == nil {
|
|
current.Summary.LongestStreakDays = previous.Summary.LongestStreakDays
|
|
}
|
|
}
|
|
if !usageAvailable {
|
|
current.Usage = append([]UsagePoint(nil), previous.Usage...)
|
|
if current.Usage == nil {
|
|
current.Usage = []UsagePoint{}
|
|
}
|
|
}
|
|
current.Summary.PeakDailyTokens = nil
|
|
if usageAvailable || len(current.Usage) > 0 {
|
|
current.CurrentCycle = currentTokenCycle(current.Limits, current.Usage, current.FetchedAt)
|
|
current.Summary.PeakDailyTokens = peakDailyTokensForCycle(current.CurrentCycle, current.Usage, current.FetchedAt)
|
|
}
|
|
}
|
|
|
|
func dashboardFromSnapshot(account store.Account, snapshot cliproxy.Snapshot) Dashboard {
|
|
authMode := "cliproxyapi"
|
|
dashboard := Dashboard{
|
|
AccountID: account.ID,
|
|
DisplayName: account.DisplayName,
|
|
Account: AccountView{
|
|
Email: snapshot.Auth.Email,
|
|
AuthMode: &authMode,
|
|
PlanType: snapshot.Auth.PlanType,
|
|
Connected: true,
|
|
},
|
|
Limits: make([]LimitBucket, 0, len(snapshot.Limits)),
|
|
Usage: make([]UsagePoint, 0, len(snapshot.Usage)),
|
|
FetchedAt: snapshot.FetchedAt.Unix(),
|
|
Stale: false,
|
|
}
|
|
for _, limit := range snapshot.Limits {
|
|
dashboard.Limits = append(dashboard.Limits, LimitBucket{
|
|
LimitID: limit.LimitID,
|
|
LimitName: limit.LimitName,
|
|
WindowType: limit.WindowType,
|
|
UsedPercent: limit.UsedPercent,
|
|
WindowDurationMinutes: limit.WindowDurationMinutes,
|
|
ResetsAt: limit.ResetsAt,
|
|
PlanType: limit.PlanType,
|
|
})
|
|
}
|
|
dashboard.Summary = UsageSummary{
|
|
LifetimeTokens: snapshot.Summary.LifetimeTokens,
|
|
PeakDailyTokens: snapshot.Summary.PeakDailyTokens,
|
|
LongestRunningTurnSec: snapshot.Summary.LongestRunningTurnSec,
|
|
CurrentStreakDays: snapshot.Summary.CurrentStreakDays,
|
|
LongestStreakDays: snapshot.Summary.LongestStreakDays,
|
|
}
|
|
for _, point := range snapshot.Usage {
|
|
dashboard.Usage = append(dashboard.Usage, UsagePoint{Date: point.Date, TotalTokens: point.TotalTokens})
|
|
}
|
|
dashboard.Summary.PeakDailyTokens = nil
|
|
if snapshot.UsageAvailable {
|
|
dashboard.CurrentCycle = currentTokenCycle(dashboard.Limits, dashboard.Usage, dashboard.FetchedAt)
|
|
dashboard.Summary.PeakDailyTokens = peakDailyTokensForCycle(dashboard.CurrentCycle, dashboard.Usage, dashboard.FetchedAt)
|
|
}
|
|
if snapshot.ResetCredits != nil && snapshot.ResetCredits.AvailableCount > 0 {
|
|
expiresAt := append([]int64(nil), snapshot.ResetCredits.ExpiresAt...)
|
|
sort.Slice(expiresAt, func(i, j int) bool { return expiresAt[i] < expiresAt[j] })
|
|
dashboard.ResetCredits = &ResetCreditsSummary{AvailableCount: snapshot.ResetCredits.AvailableCount, ExpiresAt: expiresAt}
|
|
}
|
|
return dashboard
|
|
}
|
|
|
|
func (a *App) persistDashboard(d Dashboard) (bool, error) {
|
|
tx, err := a.store.DB.Begin()
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
resetDetected, err := a.persistDashboardTx(tx, d)
|
|
if err != nil {
|
|
_ = tx.Rollback()
|
|
return false, err
|
|
}
|
|
if err = tx.Commit(); err != nil {
|
|
return false, err
|
|
}
|
|
return resetDetected, nil
|
|
}
|
|
|
|
func (a *App) persistDashboardTx(tx *sql.Tx, d Dashboard) (bool, error) {
|
|
for _, point := range d.Usage {
|
|
if _, err := tx.Exec("INSERT INTO daily_usage(account_id,date,total_tokens,fetched_at) VALUES(?,?,?,?) ON CONFLICT(account_id,date) DO UPDATE SET total_tokens=excluded.total_tokens,fetched_at=excluded.fetched_at", d.AccountID, point.Date, point.TotalTokens, d.FetchedAt); err != nil {
|
|
return false, err
|
|
}
|
|
}
|
|
resetDetected, err := a.storeLimitSnapshotsTx(tx, d)
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
result, err := tx.Exec("UPDATE accounts SET email=?,plan_type=?,connected=1,updated_at=? WHERE id=?", d.Account.Email, d.Account.PlanType, time.Now().Unix(), d.AccountID)
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
if affected, affectedErr := result.RowsAffected(); affectedErr != nil {
|
|
return false, affectedErr
|
|
} else if affected == 0 {
|
|
return false, sql.ErrNoRows
|
|
}
|
|
return resetDetected, nil
|
|
}
|
|
|
|
const resetDropTolerance = 0.01
|
|
|
|
func (a *App) storeLimitSnapshots(d Dashboard) (bool, error) {
|
|
tx, err := a.store.DB.Begin()
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
resetDetected, err := a.storeLimitSnapshotsTx(tx, d)
|
|
if err != nil {
|
|
_ = tx.Rollback()
|
|
return false, err
|
|
}
|
|
if err = tx.Commit(); err != nil {
|
|
return false, err
|
|
}
|
|
return resetDetected, nil
|
|
}
|
|
|
|
func (a *App) storeLimitSnapshotsTx(tx *sql.Tx, d Dashboard) (bool, error) {
|
|
g := a.general()
|
|
resetDetected := false
|
|
for _, limit := range d.Limits {
|
|
var previousID, previousFetchedAt, previousResetsAt int64
|
|
var previousUsed float64
|
|
err := tx.QueryRow(`SELECT id,used_percent,resets_at,fetched_at FROM limit_snapshots
|
|
WHERE account_id=? AND limit_id=? AND window_type=? ORDER BY fetched_at DESC,id DESC LIMIT 1`,
|
|
d.AccountID, limit.LimitID, limit.WindowType).Scan(&previousID, &previousUsed, &previousResetsAt, &previousFetchedAt)
|
|
if err != nil && err != sql.ErrNoRows {
|
|
return false, err
|
|
}
|
|
age := d.FetchedAt - previousFetchedAt
|
|
if err == nil && g.NotifyAfter && age >= 0 && age <= int64((6*time.Hour).Seconds()) && previousUsed-limit.UsedPercent > resetDropTolerance {
|
|
kind := "detected_after"
|
|
key := fmt.Sprintf("%d:%s:%s:detected:%d", d.AccountID, limit.LimitID, limit.WindowType, previousID)
|
|
now := time.Unix(d.FetchedAt, 0)
|
|
if previousResetsAt <= d.FetchedAt && now.Sub(time.Unix(previousResetsAt, 0)) <= 6*time.Hour {
|
|
kind = "after"
|
|
key = fmt.Sprintf("%d:%s:%s:%d:after", d.AccountID, limit.LimitID, limit.WindowType, previousResetsAt)
|
|
}
|
|
event := notificationEvent{Version: 1, Kind: kind, Account: d.DisplayName, DurationMins: limit.WindowDurationMinutes,
|
|
Remaining: 100 - limit.UsedPercent, PreviousUsed: previousUsed, Used: limit.UsedPercent, ResetsAt: limit.ResetsAt}
|
|
body, _ := json.Marshal(event)
|
|
if _, err = tx.Exec(`INSERT OR IGNORE INTO notifications
|
|
(dedupe_key,channel,kind,status,attempts,last_error,scheduled_at,sent_at,body)
|
|
VALUES(?,?,?,'pending',0,'',?,NULL,?)`, key, "configured", kind, d.FetchedAt, string(body)); err != nil {
|
|
return false, err
|
|
}
|
|
resetDetected = true
|
|
}
|
|
if _, err = tx.Exec("INSERT INTO limit_snapshots(limit_id,window_type,used_percent,duration_mins,resets_at,fetched_at,account_id) VALUES(?,?,?,?,?,?,?)", limit.LimitID, limit.WindowType, limit.UsedPercent, limit.WindowDurationMinutes, limit.ResetsAt, d.FetchedAt, d.AccountID); err != nil {
|
|
return false, err
|
|
}
|
|
}
|
|
return resetDetected, nil
|
|
}
|
|
|
|
func currentTokenCycle(limits []LimitBucket, usage []UsagePoint, fetchedAt int64) *TokenCycle {
|
|
current := LimitBucket{}
|
|
found := false
|
|
for _, limit := range limits {
|
|
if limit.WindowDurationMinutes <= 0 || limit.ResetsAt <= fetchedAt {
|
|
continue
|
|
}
|
|
if !found || betterTokenCycleLimit(limit, current) {
|
|
current = limit
|
|
found = true
|
|
}
|
|
}
|
|
if !found {
|
|
return nil
|
|
}
|
|
startedAt := current.ResetsAt - int64(current.WindowDurationMinutes)*60
|
|
startDate := time.Unix(startedAt, 0).UTC().Format("2006-01-02")
|
|
endDate := time.Unix(fetchedAt, 0).UTC().Format("2006-01-02")
|
|
var total int64
|
|
for _, point := range usage {
|
|
if point.Date < startDate || point.Date > endDate {
|
|
continue
|
|
}
|
|
if _, err := time.Parse("2006-01-02", point.Date); err != nil {
|
|
continue
|
|
}
|
|
total += point.TotalTokens
|
|
}
|
|
return &TokenCycle{LimitID: current.LimitID, WindowType: current.WindowType, WindowDurationMinutes: current.WindowDurationMinutes, StartedAt: startedAt, ResetsAt: current.ResetsAt, TotalTokens: total}
|
|
}
|
|
|
|
func peakDailyTokensForCycle(cycle *TokenCycle, usage []UsagePoint, fetchedAt int64) *int64 {
|
|
if cycle == nil {
|
|
return nil
|
|
}
|
|
startDate := time.Unix(cycle.StartedAt, 0).UTC().Format("2006-01-02")
|
|
endDate := time.Unix(fetchedAt, 0).UTC().Format("2006-01-02")
|
|
var peak int64
|
|
found := false
|
|
for _, point := range usage {
|
|
if _, err := time.Parse("2006-01-02", point.Date); err != nil || point.Date < startDate || point.Date > endDate {
|
|
continue
|
|
}
|
|
if !found || point.TotalTokens > peak {
|
|
peak = point.TotalTokens
|
|
found = true
|
|
}
|
|
}
|
|
if !found {
|
|
return nil
|
|
}
|
|
return &peak
|
|
}
|
|
|
|
func betterTokenCycleLimit(candidate, current LimitBucket) bool {
|
|
if candidate.WindowDurationMinutes != current.WindowDurationMinutes {
|
|
return candidate.WindowDurationMinutes > current.WindowDurationMinutes
|
|
}
|
|
if candidate.WindowType != current.WindowType {
|
|
return candidate.WindowType == "secondary"
|
|
}
|
|
if candidate.LimitID != current.LimitID {
|
|
return candidate.LimitID < current.LimitID
|
|
}
|
|
return candidate.ResetsAt > current.ResetsAt
|
|
}
|