This commit is contained in:
@@ -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