feat: validate Team workspace connections
This commit is contained in:
@@ -13,6 +13,7 @@ import (
|
||||
"time"
|
||||
|
||||
"codex-helper/internal/security"
|
||||
"codex-helper/internal/store"
|
||||
)
|
||||
|
||||
func (a *App) api(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -57,7 +58,8 @@ func (a *App) api(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
case p == "accounts" && r.Method == "POST":
|
||||
var in struct {
|
||||
DisplayName string `json:"displayName"`
|
||||
DisplayName string `json:"displayName"`
|
||||
ExpectedKind string `json:"expectedKind"`
|
||||
}
|
||||
if decode(r, &in) != nil {
|
||||
jsonOut(w, 400, map[string]string{"error": "请求格式错误"})
|
||||
@@ -67,7 +69,14 @@ func (a *App) api(w http.ResponseWriter, r *http.Request) {
|
||||
if in.DisplayName == "" {
|
||||
in.DisplayName = "新账号"
|
||||
}
|
||||
x, e := a.store.CreateAccount(in.DisplayName)
|
||||
if in.ExpectedKind == "" {
|
||||
in.ExpectedKind = "any"
|
||||
}
|
||||
if !store.ValidExpectedKind(in.ExpectedKind) {
|
||||
jsonOut(w, 400, map[string]string{"error": "连接类型无效"})
|
||||
break
|
||||
}
|
||||
x, e := a.store.CreateAccount(in.DisplayName, in.ExpectedKind)
|
||||
if e == nil {
|
||||
a.addRuntime(x.ID)
|
||||
jsonOut(w, 201, x)
|
||||
@@ -164,13 +173,27 @@ func (a *App) accountAPI(w http.ResponseWriter, r *http.Request, p string) {
|
||||
switch {
|
||||
case action == "" && r.Method == "PUT":
|
||||
var in struct {
|
||||
DisplayName string `json:"displayName"`
|
||||
DisplayName string `json:"displayName"`
|
||||
ExpectedKind string `json:"expectedKind"`
|
||||
}
|
||||
if decode(r, &in) != nil || strings.TrimSpace(in.DisplayName) == "" {
|
||||
jsonOut(w, 400, map[string]string{"error": "名称不能为空"})
|
||||
return
|
||||
}
|
||||
e = a.store.RenameAccount(id, strings.TrimSpace(in.DisplayName))
|
||||
if in.ExpectedKind == "" {
|
||||
accounts, _ := a.store.Accounts()
|
||||
for _, account := range accounts {
|
||||
if account.ID == id {
|
||||
in.ExpectedKind = account.ExpectedKind
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
if !store.ValidExpectedKind(in.ExpectedKind) {
|
||||
jsonOut(w, 400, map[string]string{"error": "连接类型无效"})
|
||||
return
|
||||
}
|
||||
e = a.store.UpdateAccountSettings(id, strings.TrimSpace(in.DisplayName), in.ExpectedKind)
|
||||
if e == nil {
|
||||
jsonOut(w, 200, map[string]bool{"ok": true})
|
||||
}
|
||||
@@ -393,6 +416,24 @@ func (a *App) syncAccount(ctx context.Context, id int64) error {
|
||||
d.Limits = flattenLimit(*lr.RateLimits)
|
||||
}
|
||||
}
|
||||
// Some app-server versions expose the workspace plan only on limit buckets.
|
||||
if ar.Account != nil && store.AccountKind(d.Account.PlanType) == "unknown" {
|
||||
var fallback *string
|
||||
consistent := true
|
||||
for _, x := range d.Limits {
|
||||
if store.AccountKind(x.PlanType) == "unknown" {
|
||||
continue
|
||||
}
|
||||
if fallback == nil {
|
||||
fallback = x.PlanType
|
||||
} else if store.AccountKind(fallback) != store.AccountKind(x.PlanType) {
|
||||
consistent = false
|
||||
}
|
||||
}
|
||||
if consistent && fallback != nil {
|
||||
d.Account.PlanType = fallback
|
||||
}
|
||||
}
|
||||
var ur struct {
|
||||
Summary UsageSummary `json:"summary"`
|
||||
Daily []struct {
|
||||
|
||||
@@ -135,11 +135,45 @@ func (a *App) keepCodex() {
|
||||
}
|
||||
func (a *App) onCodexNotification(id int64) func(string, json.RawMessage) {
|
||||
return func(method string, _ json.RawMessage) {
|
||||
if method == "account/updated" || method == "account/rateLimits/updated" {
|
||||
go a.syncAccount(context.Background(), id)
|
||||
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 {
|
||||
|
||||
@@ -158,3 +158,26 @@ func TestDeviceLoginStartsColdRuntime(t *testing.T) {
|
||||
t.Fatalf("starts = %d, initializes = %d, calls = %d; want 1 each", starts, initializes, calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAccountClassificationReadyRequiresConnectedKnownPlan(t *testing.T) {
|
||||
team := "team"
|
||||
unknown := "unknown"
|
||||
tests := []struct {
|
||||
name string
|
||||
account AccountView
|
||||
want bool
|
||||
}{
|
||||
{name: "disconnected", account: AccountView{PlanType: &team}, want: false},
|
||||
{name: "missing plan", account: AccountView{Connected: true}, want: false},
|
||||
{name: "unknown plan", account: AccountView{Connected: true, PlanType: &unknown}, want: false},
|
||||
{name: "classified", account: AccountView{Connected: true, PlanType: &team}, want: true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
rt := &accountRuntime{dash: Dashboard{Account: tt.account}}
|
||||
if got := accountClassificationReady(rt); got != tt.want {
|
||||
t.Fatalf("accountClassificationReady() = %v; want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
@@ -59,12 +60,31 @@ func (s *Store) migrateAccounts() error {
|
||||
display_name TEXT NOT NULL,
|
||||
email TEXT,
|
||||
plan_type TEXT,
|
||||
expected_kind TEXT NOT NULL DEFAULT 'any',
|
||||
connected INTEGER NOT NULL DEFAULT 0,
|
||||
created_at INTEGER NOT NULL,
|
||||
updated_at INTEGER NOT NULL
|
||||
)`); err != nil {
|
||||
return err
|
||||
}
|
||||
var hasExpectedKind bool
|
||||
rows, qerr := tx.Query("PRAGMA table_info(accounts)")
|
||||
if qerr != nil {
|
||||
return qerr
|
||||
}
|
||||
for rows.Next() {
|
||||
var cid, notnull, pk int
|
||||
var name, typ string
|
||||
var def any
|
||||
_ = rows.Scan(&cid, &name, &typ, ¬null, &def, &pk)
|
||||
hasExpectedKind = hasExpectedKind || name == "expected_kind"
|
||||
}
|
||||
rows.Close()
|
||||
if !hasExpectedKind {
|
||||
if _, err = tx.Exec("ALTER TABLE accounts ADD COLUMN expected_kind TEXT NOT NULL DEFAULT 'any'"); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
var count int
|
||||
if err = tx.QueryRow("SELECT COUNT(*) FROM accounts").Scan(&count); err != nil {
|
||||
return err
|
||||
@@ -122,17 +142,21 @@ func (s *Store) migrateAccounts() error {
|
||||
}
|
||||
|
||||
type Account struct {
|
||||
ID int64 `json:"id"`
|
||||
DisplayName string `json:"displayName"`
|
||||
Email *string `json:"email"`
|
||||
PlanType *string `json:"planType"`
|
||||
Connected bool `json:"connected"`
|
||||
CreatedAt int64 `json:"createdAt"`
|
||||
UpdatedAt int64 `json:"updatedAt"`
|
||||
ID int64 `json:"id"`
|
||||
DisplayName string `json:"displayName"`
|
||||
Email *string `json:"email"`
|
||||
PlanType *string `json:"planType"`
|
||||
ExpectedKind string `json:"expectedKind"`
|
||||
ActualKind string `json:"actualKind"`
|
||||
ValidationStatus string `json:"validationStatus"`
|
||||
PossibleDuplicate bool `json:"possibleDuplicate"`
|
||||
Connected bool `json:"connected"`
|
||||
CreatedAt int64 `json:"createdAt"`
|
||||
UpdatedAt int64 `json:"updatedAt"`
|
||||
}
|
||||
|
||||
func (s *Store) Accounts() ([]Account, error) {
|
||||
rows, e := s.DB.Query("SELECT id,display_name,email,plan_type,connected,created_at,updated_at FROM accounts ORDER BY id")
|
||||
rows, e := s.DB.Query("SELECT id,display_name,email,plan_type,expected_kind,connected,created_at,updated_at FROM accounts ORDER BY id")
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
@@ -140,24 +164,40 @@ 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.Connected, &a.CreatedAt, &a.UpdatedAt); e != nil {
|
||||
if e = rows.Scan(&a.ID, &a.DisplayName, &a.Email, &a.PlanType, &a.ExpectedKind, &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)
|
||||
out = append(out, a)
|
||||
}
|
||||
for i := range out {
|
||||
if out[i].Email == nil || out[i].ActualKind == "unknown" {
|
||||
continue
|
||||
}
|
||||
for j := range out {
|
||||
if i != j && out[j].Email != nil && strings.EqualFold(*out[i].Email, *out[j].Email) && out[i].ActualKind == out[j].ActualKind {
|
||||
out[i].PossibleDuplicate = true
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
func (s *Store) CreateAccount(name string) (Account, error) {
|
||||
func (s *Store) CreateAccount(name string, 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,created_at,updated_at) VALUES(?,?,?)", name, now, now)
|
||||
r, e := s.DB.Exec("INSERT INTO accounts(display_name,expected_kind,created_at,updated_at) VALUES(?,?,?,?)", name, expectedKind, now, now)
|
||||
if e != nil {
|
||||
return Account{}, e
|
||||
}
|
||||
id, _ := r.LastInsertId()
|
||||
return Account{ID: id, DisplayName: name, CreatedAt: now, UpdatedAt: now}, nil
|
||||
return Account{ID: id, DisplayName: name, ExpectedKind: expectedKind, ActualKind: "unknown", ValidationStatus: "pending", CreatedAt: now, UpdatedAt: now}, nil
|
||||
}
|
||||
func (s *Store) RenameAccount(id int64, name string) error {
|
||||
r, e := s.DB.Exec("UPDATE accounts SET display_name=?,updated_at=? WHERE id=?", name, time.Now().Unix(), id)
|
||||
func (s *Store) UpdateAccountSettings(id int64, name, expectedKind string) error {
|
||||
r, e := s.DB.Exec("UPDATE accounts SET display_name=?,expected_kind=?,updated_at=? WHERE id=?", name, expectedKind, time.Now().Unix(), id)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
@@ -167,6 +207,45 @@ func (s *Store) RenameAccount(id int64, name string) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
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 {
|
||||
return err
|
||||
}
|
||||
return s.UpdateAccountSettings(id, name, kind)
|
||||
}
|
||||
|
||||
func ValidExpectedKind(kind string) bool {
|
||||
return kind == "any" || kind == "personal" || kind == "team"
|
||||
}
|
||||
|
||||
func AccountKind(plan *string) string {
|
||||
if plan == nil {
|
||||
return "unknown"
|
||||
}
|
||||
switch strings.ToLower(strings.TrimSpace(*plan)) {
|
||||
case "free", "go", "plus", "pro", "prolite":
|
||||
return "personal"
|
||||
case "team", "business", "self_serve_business_prolite", "self_serve_business_usage_based":
|
||||
return "team"
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
}
|
||||
|
||||
func validationStatus(expected string, connected bool, plan *string) string {
|
||||
if !connected {
|
||||
return "pending"
|
||||
}
|
||||
actual := AccountKind(plan)
|
||||
if actual == "unknown" {
|
||||
return "unknown"
|
||||
}
|
||||
if expected == "any" || expected == actual {
|
||||
return "matched"
|
||||
}
|
||||
return "mismatch"
|
||||
}
|
||||
func (s *Store) UpdateAccount(id int64, email, plan *string, connected bool) error {
|
||||
_, e := s.DB.Exec("UPDATE accounts SET email=?,plan_type=?,connected=?,updated_at=? WHERE id=?", email, plan, connected, time.Now().Unix(), id)
|
||||
return e
|
||||
|
||||
@@ -54,6 +54,56 @@ func TestAccountsAndPerAccountUsage(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestAccountKindAndValidation(t *testing.T) {
|
||||
tests := []struct{ plan, kind string }{
|
||||
{"plus", "personal"}, {"Pro", "personal"}, {"team", "team"},
|
||||
{"business", "team"}, {"self_serve_business_usage_based", "team"},
|
||||
{"enterprise", "unknown"}, {"", "unknown"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
plan := tt.plan
|
||||
if got := AccountKind(&plan); got != tt.kind {
|
||||
t.Errorf("AccountKind(%q) = %q; want %q", tt.plan, got, tt.kind)
|
||||
}
|
||||
}
|
||||
if got := validationStatus("team", true, ptr("plus")); got != "mismatch" {
|
||||
t.Fatalf("team/plus validation = %q; want mismatch", got)
|
||||
}
|
||||
if got := validationStatus("team", true, ptr("business")); got != "matched" {
|
||||
t.Fatalf("team/business validation = %q; want matched", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExistingAccountsGainExpectedKind(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, connected INTEGER NOT NULL DEFAULT 0, created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL
|
||||
); INSERT INTO accounts VALUES(1,'旧连接','user@example.com','team',1,1,1);`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db.Close()
|
||||
s, err := Open(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer s.DB.Close()
|
||||
accounts, err := s.Accounts()
|
||||
if err != nil || len(accounts) != 1 {
|
||||
t.Fatalf("accounts = %#v, %v", accounts, err)
|
||||
}
|
||||
if accounts[0].ExpectedKind != "any" || accounts[0].ActualKind != "team" || accounts[0].ValidationStatus != "matched" {
|
||||
t.Fatalf("migrated account = %#v", accounts[0])
|
||||
}
|
||||
}
|
||||
|
||||
func ptr(value string) *string { return &value }
|
||||
|
||||
func TestLegacyUsageMigratesToDefaultAccount(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
db, err := sql.Open("sqlite", filepath.Join(dir, "codex-helper.db"))
|
||||
|
||||
Vendored
+2
-2
@@ -1,3 +1,3 @@
|
||||
<script type="module" crossorigin src="/assets/index-DhESEhcz.js"></script>
|
||||
<link rel="stylesheet" crossorigin href="/assets/index-BpH_Y97H.css">
|
||||
<script type="module" crossorigin src="/assets/index-BbkhVAJw.js"></script>
|
||||
<link rel="stylesheet" crossorigin href="/assets/index-BTO6fDiH.css">
|
||||
<div id="root"></div>
|
||||
|
||||
Reference in New Issue
Block a user