583 lines
21 KiB
Go
583 lines
21 KiB
Go
package app
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"reflect"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"codex-helper/internal/security"
|
|
"codex-helper/internal/store"
|
|
)
|
|
|
|
func TestSystemStatusRejectsNonGETMethods(t *testing.T) {
|
|
a := newReminderTestApp(t)
|
|
recorder := httptest.NewRecorder()
|
|
request := httptest.NewRequest(http.MethodPost, "/api/v1/system/status", nil)
|
|
a.api(recorder, request)
|
|
if recorder.Code != http.StatusMethodNotAllowed {
|
|
t.Fatalf("status = %d, body = %s", recorder.Code, recorder.Body.String())
|
|
}
|
|
}
|
|
|
|
func TestSystemStatusReturnsBuildVersion(t *testing.T) {
|
|
a := newReminderTestApp(t)
|
|
originalVersion := Version
|
|
Version = "1.2.3-test"
|
|
t.Cleanup(func() { Version = originalVersion })
|
|
recorder := httptest.NewRecorder()
|
|
request := httptest.NewRequest(http.MethodGet, "/api/v1/system/status", nil)
|
|
a.api(recorder, request)
|
|
if recorder.Code != http.StatusOK {
|
|
t.Fatalf("status = %d, body = %s", recorder.Code, recorder.Body.String())
|
|
}
|
|
var body struct {
|
|
Version string `json:"version"`
|
|
}
|
|
if err := json.Unmarshal(recorder.Body.Bytes(), &body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if body.Version != "1.2.3-test" {
|
|
t.Fatalf("version = %q; want %q", body.Version, "1.2.3-test")
|
|
}
|
|
}
|
|
|
|
func TestDashboardSerializesNilListsAsEmptyArrays(t *testing.T) {
|
|
a := newReminderTestApp(t)
|
|
if err := a.store.Set("initialized", "true"); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
a.runtimes[1] = &accountRuntime{}
|
|
_, err := a.store.DB.Exec("INSERT INTO sessions(token_hash,expires_at,created_at) VALUES(?,?,?)", security.HashToken("test-session"), time.Now().Add(time.Hour).Unix(), time.Now().Unix())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
recorder := httptest.NewRecorder()
|
|
request := httptest.NewRequest(http.MethodGet, "/api/v1/dashboard?accountId=1", nil)
|
|
request.AddCookie(&http.Cookie{Name: "session", Value: "test-session"})
|
|
a.api(recorder, request)
|
|
if recorder.Code != http.StatusOK {
|
|
t.Fatalf("status = %d, body = %s", recorder.Code, recorder.Body.String())
|
|
}
|
|
var body struct {
|
|
Limits []LimitBucket `json:"limits"`
|
|
Usage []UsagePoint `json:"usage"`
|
|
}
|
|
if err := json.Unmarshal(recorder.Body.Bytes(), &body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if body.Limits == nil || body.Usage == nil {
|
|
t.Fatalf("nil lists in response: %s", recorder.Body.String())
|
|
}
|
|
}
|
|
|
|
func TestAnonymousOverviewIsReadOnly(t *testing.T) {
|
|
a := newReminderTestApp(t)
|
|
if err := a.store.Set("initialized", "true"); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
email := "owner@example.com"
|
|
plan := "plus"
|
|
if err := a.store.UpdateAccount(1, &email, &plan, true); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
publicVisible := true
|
|
if err := a.store.UpdateAccountSettingsWithVisibility(1, "默认账号", "any", &publicVisible); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
a.runtimes[1] = &accountRuntime{
|
|
dash: Dashboard{
|
|
AccountID: 1,
|
|
DisplayName: "默认账号",
|
|
Account: AccountView{Email: &email, PlanType: &plan, Connected: true},
|
|
Limits: []LimitBucket{},
|
|
Usage: []UsagePoint{},
|
|
ResetCredits: &ResetCreditsSummary{AvailableCount: 2, ExpiresAt: []int64{1784246400}},
|
|
FetchedAt: time.Now().Unix(),
|
|
},
|
|
}
|
|
|
|
accountsRecorder := httptest.NewRecorder()
|
|
a.api(accountsRecorder, httptest.NewRequest(http.MethodGet, "/api/v1/accounts", nil))
|
|
if accountsRecorder.Code != http.StatusOK {
|
|
t.Fatalf("anonymous accounts status = %d, body = %s", accountsRecorder.Code, accountsRecorder.Body.String())
|
|
}
|
|
var accounts []struct {
|
|
Email *string `json:"email"`
|
|
ExpectedKind string `json:"expectedKind"`
|
|
PublicVisible bool `json:"publicVisible"`
|
|
ValidationState string `json:"validationStatus"`
|
|
}
|
|
if err := json.Unmarshal(accountsRecorder.Body.Bytes(), &accounts); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if len(accounts) != 1 || accounts[0].Email != nil || !accounts[0].PublicVisible || accounts[0].ExpectedKind != "any" || accounts[0].ValidationState != "unknown" {
|
|
t.Fatalf("anonymous account data = %#v; sensitive account fields were not redacted", accounts)
|
|
}
|
|
|
|
dashboardRecorder := httptest.NewRecorder()
|
|
a.api(dashboardRecorder, httptest.NewRequest(http.MethodGet, "/api/v1/dashboard?accountId=1", nil))
|
|
if dashboardRecorder.Code != http.StatusOK {
|
|
t.Fatalf("anonymous dashboard status = %d, body = %s", dashboardRecorder.Code, dashboardRecorder.Body.String())
|
|
}
|
|
var publicDashboardBody Dashboard
|
|
if err := json.Unmarshal(dashboardRecorder.Body.Bytes(), &publicDashboardBody); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if publicDashboardBody.Account.Email != nil || publicDashboardBody.Account.AuthMode != nil {
|
|
t.Fatalf("anonymous dashboard account = %#v; identity fields were not redacted", publicDashboardBody.Account)
|
|
}
|
|
if publicDashboardBody.ResetCredits == nil || publicDashboardBody.ResetCredits.AvailableCount != 2 {
|
|
t.Fatalf("anonymous dashboard reset credits = %#v; want two credits", publicDashboardBody.ResetCredits)
|
|
}
|
|
|
|
for _, path := range []string{"/api/v1/settings/general", "/api/v1/accounts", "/api/v1/accounts/1/sync"} {
|
|
recorder := httptest.NewRecorder()
|
|
method := http.MethodGet
|
|
if path == "/api/v1/accounts" || strings.HasSuffix(path, "/sync") {
|
|
method = http.MethodPost
|
|
}
|
|
a.api(recorder, httptest.NewRequest(method, path, nil))
|
|
if recorder.Code != http.StatusUnauthorized {
|
|
t.Fatalf("anonymous %s status = %d, body = %s; configuration must require login", path, recorder.Code, recorder.Body.String())
|
|
}
|
|
}
|
|
|
|
session := "test-session"
|
|
if _, err := a.store.DB.Exec("INSERT INTO sessions(token_hash,expires_at,created_at) VALUES(?,?,?)", security.HashToken(session), time.Now().Add(time.Hour).Unix(), time.Now().Unix()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
privateRecorder := httptest.NewRecorder()
|
|
privateRequest := httptest.NewRequest(http.MethodGet, "/api/v1/dashboard?accountId=1", nil)
|
|
privateRequest.AddCookie(&http.Cookie{Name: "session", Value: session})
|
|
a.api(privateRecorder, privateRequest)
|
|
if privateRecorder.Code != http.StatusOK {
|
|
t.Fatalf("authenticated dashboard status = %d, body = %s", privateRecorder.Code, privateRecorder.Body.String())
|
|
}
|
|
var privateDashboardBody Dashboard
|
|
if err := json.Unmarshal(privateRecorder.Body.Bytes(), &privateDashboardBody); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if privateDashboardBody.Account.Email == nil || *privateDashboardBody.Account.Email != email {
|
|
t.Fatalf("authenticated dashboard email = %v; want %q", privateDashboardBody.Account.Email, email)
|
|
}
|
|
if privateDashboardBody.ResetCredits == nil || privateDashboardBody.ResetCredits.AvailableCount != 2 {
|
|
t.Fatalf("authenticated dashboard reset credits = %#v; want two credits", privateDashboardBody.ResetCredits)
|
|
}
|
|
configRecorder := httptest.NewRecorder()
|
|
configRequest := httptest.NewRequest(http.MethodPut, "/api/v1/settings/general", strings.NewReader(`{"timezone":"UTC","theme":"system","syncMinutes":5,"retentionDays":90,"beforeMinutes":30,"notifyBefore":true,"notifyAfter":true}`))
|
|
configRequest.AddCookie(&http.Cookie{Name: "session", Value: session})
|
|
configRequest.Header.Set("X-Requested-With", "codex-helper")
|
|
a.api(configRecorder, configRequest)
|
|
if configRecorder.Code != http.StatusOK {
|
|
t.Fatalf("authenticated settings status = %d, body = %s", configRecorder.Code, configRecorder.Body.String())
|
|
}
|
|
}
|
|
|
|
func TestAccountVisibilityFiltersAnonymousOverviewAndCanBeUpdated(t *testing.T) {
|
|
a := newReminderTestApp(t)
|
|
if err := a.store.Set("initialized", "true"); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
publicAccount, err := a.store.CreateAccountWithVisibility("公开账号", "team", true)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
privateAccount, err := a.store.CreateAccount("私有账号", "personal")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
a.runtimes[publicAccount.ID] = &accountRuntime{}
|
|
a.runtimes[privateAccount.ID] = &accountRuntime{}
|
|
|
|
accountsRecorder := httptest.NewRecorder()
|
|
a.api(accountsRecorder, httptest.NewRequest(http.MethodGet, "/api/v1/accounts", nil))
|
|
if accountsRecorder.Code != http.StatusOK {
|
|
t.Fatalf("anonymous accounts status = %d, body = %s", accountsRecorder.Code, accountsRecorder.Body.String())
|
|
}
|
|
var visible []struct {
|
|
ID int64 `json:"id"`
|
|
PublicVisible bool `json:"publicVisible"`
|
|
}
|
|
if err := json.Unmarshal(accountsRecorder.Body.Bytes(), &visible); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if len(visible) != 1 || visible[0].ID != publicAccount.ID || !visible[0].PublicVisible {
|
|
t.Fatalf("anonymous accounts = %#v; want only public account %d", visible, publicAccount.ID)
|
|
}
|
|
|
|
privateDashboard := httptest.NewRecorder()
|
|
privatePath := "/api/v1/dashboard?accountId=" + strconv.FormatInt(privateAccount.ID, 10)
|
|
a.api(privateDashboard, httptest.NewRequest(http.MethodGet, privatePath, nil))
|
|
if privateDashboard.Code != http.StatusNotFound {
|
|
t.Fatalf("anonymous private dashboard status = %d, body = %s", privateDashboard.Code, privateDashboard.Body.String())
|
|
}
|
|
|
|
publicDashboard := httptest.NewRecorder()
|
|
publicPath := "/api/v1/dashboard?accountId=" + strconv.FormatInt(publicAccount.ID, 10)
|
|
a.api(publicDashboard, httptest.NewRequest(http.MethodGet, publicPath, nil))
|
|
if publicDashboard.Code != http.StatusOK {
|
|
t.Fatalf("anonymous public dashboard status = %d, body = %s", publicDashboard.Code, publicDashboard.Body.String())
|
|
}
|
|
|
|
session := "visibility-session"
|
|
if _, err := a.store.DB.Exec("INSERT INTO sessions(token_hash,expires_at,created_at) VALUES(?,?,?)", security.HashToken(session), time.Now().Add(time.Hour).Unix(), time.Now().Unix()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
createRecorder := httptest.NewRecorder()
|
|
createRequest := httptest.NewRequest(http.MethodPost, "/api/v1/accounts", strings.NewReader(`{"displayName":"接口公开账号","expectedKind":"team","publicVisible":true}`))
|
|
createRequest.AddCookie(&http.Cookie{Name: "session", Value: session})
|
|
createRequest.Header.Set("X-Requested-With", "codex-helper")
|
|
a.api(createRecorder, createRequest)
|
|
if createRecorder.Code != http.StatusCreated {
|
|
t.Fatalf("authenticated account creation status = %d, body = %s", createRecorder.Code, createRecorder.Body.String())
|
|
}
|
|
var created store.Account
|
|
if err := json.Unmarshal(createRecorder.Body.Bytes(), &created); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if !created.PublicVisible {
|
|
t.Fatalf("created account = %#v; want publicVisible=true", created)
|
|
}
|
|
authenticatedAccounts := httptest.NewRecorder()
|
|
authenticatedRequest := httptest.NewRequest(http.MethodGet, "/api/v1/accounts", nil)
|
|
authenticatedRequest.AddCookie(&http.Cookie{Name: "session", Value: session})
|
|
a.api(authenticatedAccounts, authenticatedRequest)
|
|
if authenticatedAccounts.Code != http.StatusOK {
|
|
t.Fatalf("authenticated accounts status = %d, body = %s", authenticatedAccounts.Code, authenticatedAccounts.Body.String())
|
|
}
|
|
var all []struct {
|
|
ID int64 `json:"id"`
|
|
}
|
|
if err := json.Unmarshal(authenticatedAccounts.Body.Bytes(), &all); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if len(all) != 4 {
|
|
t.Fatalf("authenticated accounts = %#v; want default, public, private, and newly created accounts", all)
|
|
}
|
|
|
|
anonymousUpdate := httptest.NewRecorder()
|
|
anonymousUpdateRequest := httptest.NewRequest(http.MethodPut, "/api/v1/accounts/"+strconv.FormatInt(privateAccount.ID, 10), strings.NewReader(`{"displayName":"私有账号","expectedKind":"personal","publicVisible":true}`))
|
|
a.api(anonymousUpdate, anonymousUpdateRequest)
|
|
if anonymousUpdate.Code != http.StatusUnauthorized {
|
|
t.Fatalf("anonymous visibility update status = %d, body = %s", anonymousUpdate.Code, anonymousUpdate.Body.String())
|
|
}
|
|
|
|
authenticatedUpdate := httptest.NewRecorder()
|
|
authenticatedUpdateRequest := httptest.NewRequest(http.MethodPut, "/api/v1/accounts/"+strconv.FormatInt(privateAccount.ID, 10), strings.NewReader(`{"displayName":"私有账号","expectedKind":"personal","publicVisible":true}`))
|
|
authenticatedUpdateRequest.AddCookie(&http.Cookie{Name: "session", Value: session})
|
|
authenticatedUpdateRequest.Header.Set("X-Requested-With", "codex-helper")
|
|
a.api(authenticatedUpdate, authenticatedUpdateRequest)
|
|
if authenticatedUpdate.Code != http.StatusOK {
|
|
t.Fatalf("authenticated visibility update status = %d, body = %s", authenticatedUpdate.Code, authenticatedUpdate.Body.String())
|
|
}
|
|
updated, err := a.store.Account(privateAccount.ID)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if !updated.PublicVisible {
|
|
t.Fatalf("updated account = %#v; want publicVisible=true", updated)
|
|
}
|
|
}
|
|
|
|
func TestCurrentTokenCycleUsesLongestWindowAndFiltersDailyUsage(t *testing.T) {
|
|
now := time.Date(2026, time.August, 14, 12, 0, 0, 0, time.UTC)
|
|
reset := time.Date(2026, time.August, 15, 0, 0, 0, 0, time.UTC)
|
|
cycle := currentTokenCycle(
|
|
[]LimitBucket{
|
|
{
|
|
LimitID: "codex",
|
|
WindowType: "primary",
|
|
WindowDurationMinutes: 300,
|
|
ResetsAt: now.Add(5 * time.Hour).Unix(),
|
|
},
|
|
{
|
|
LimitID: "codex",
|
|
WindowType: "secondary",
|
|
WindowDurationMinutes: 7 * 24 * 60,
|
|
ResetsAt: reset.Unix(),
|
|
},
|
|
},
|
|
[]UsagePoint{
|
|
{Date: "2026-08-07", TotalTokens: 50},
|
|
{Date: "2026-08-08", TotalTokens: 100},
|
|
{Date: "2026-08-10", TotalTokens: 200},
|
|
{Date: "2026-08-14", TotalTokens: 300},
|
|
{Date: "2026-08-15", TotalTokens: 400},
|
|
},
|
|
now.Unix(),
|
|
)
|
|
if cycle == nil {
|
|
t.Fatal("currentTokenCycle() returned nil")
|
|
}
|
|
if cycle.WindowType != "secondary" || cycle.WindowDurationMinutes != 7*24*60 {
|
|
t.Fatalf("cycle window = %#v; want the seven-day secondary window", cycle)
|
|
}
|
|
if cycle.StartedAt != time.Date(2026, time.August, 8, 0, 0, 0, 0, time.UTC).Unix() {
|
|
t.Fatalf("cycle start = %d; want 2026-08-08", cycle.StartedAt)
|
|
}
|
|
if cycle.ResetsAt != reset.Unix() || cycle.TotalTokens != 600 {
|
|
t.Fatalf("cycle = %#v; want reset %d and 600 tokens", cycle, reset.Unix())
|
|
}
|
|
}
|
|
|
|
func TestNormalizeResetCreditsKeepsAvailableCountAndUniqueExpiryTimes(t *testing.T) {
|
|
earlier := int64(1781654400)
|
|
later := int64(1784246400)
|
|
credits := normalizeResetCredits(&rawResetCredits{
|
|
AvailableCount: 3,
|
|
Credits: []rawResetCredit{
|
|
{ExpiresAt: &later},
|
|
{ExpiresAt: nil},
|
|
{ExpiresAt: &earlier},
|
|
{ExpiresAt: &later},
|
|
},
|
|
})
|
|
if credits == nil || credits.AvailableCount != 3 {
|
|
t.Fatalf("reset credits = %#v; want three available credits", credits)
|
|
}
|
|
if want := []int64{earlier, later}; !reflect.DeepEqual(credits.ExpiresAt, want) {
|
|
t.Fatalf("expiry times = %v; want %v", credits.ExpiresAt, want)
|
|
}
|
|
if normalizeResetCredits(&rawResetCredits{}) != nil {
|
|
t.Fatal("zero available credits should be hidden")
|
|
}
|
|
}
|
|
|
|
func TestCurrentTokenCycleRequiresAValidFutureResetWindow(t *testing.T) {
|
|
cycle := currentTokenCycle(
|
|
[]LimitBucket{{WindowDurationMinutes: 0, ResetsAt: 1}},
|
|
[]UsagePoint{{Date: "2026-08-14", TotalTokens: 100}},
|
|
time.Date(2026, time.August, 14, 12, 0, 0, 0, time.UTC).Unix(),
|
|
)
|
|
if cycle != nil {
|
|
t.Fatalf("currentTokenCycle() = %#v; want nil", cycle)
|
|
}
|
|
}
|
|
|
|
func TestPeakDailyTokensUsesOnlyTheCurrentTokenCycle(t *testing.T) {
|
|
now := time.Date(2026, time.August, 14, 12, 0, 0, 0, time.UTC)
|
|
cycle := &TokenCycle{
|
|
StartedAt: time.Date(2026, time.August, 8, 0, 0, 0, 0, time.UTC).Unix(),
|
|
ResetsAt: time.Date(2026, time.August, 15, 0, 0, 0, 0, time.UTC).Unix(),
|
|
}
|
|
peak := peakDailyTokensForCycle(cycle, []UsagePoint{
|
|
{Date: "2026-08-07", TotalTokens: 900},
|
|
{Date: "2026-08-08", TotalTokens: 100},
|
|
{Date: "2026-08-10", TotalTokens: 200},
|
|
{Date: "2026-08-14", TotalTokens: 300},
|
|
{Date: "2026-08-15", TotalTokens: 800},
|
|
}, now.Unix())
|
|
if peak == nil || *peak != 300 {
|
|
t.Fatalf("cycle peak = %v; want 300", peak)
|
|
}
|
|
}
|
|
|
|
func TestAutomaticSyncIntervalIsFixedAtFiveMinutes(t *testing.T) {
|
|
a := newReminderTestApp(t)
|
|
if err := a.store.SetJSON("general", GeneralSettings{SyncMinutes: 60, RetentionDays: 90, BeforeMinutes: 30}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if got := a.general().SyncMinutes; got != automaticSyncMinutes {
|
|
t.Fatalf("sync minutes = %d; want %d", got, automaticSyncMinutes)
|
|
}
|
|
if automaticSyncInterval != 5*time.Minute {
|
|
t.Fatalf("automatic sync interval = %s; want 5m", automaticSyncInterval)
|
|
}
|
|
}
|
|
|
|
func TestFlattenLimitReadsAppServerWindowDuration(t *testing.T) {
|
|
var limit rawLimit
|
|
if err := json.Unmarshal([]byte(`{
|
|
"limitId":"codex",
|
|
"primary":{"usedPercent":20,"windowDurationMins":10080,"resetsAt":1787197043}
|
|
}`), &limit); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
flattened := flattenLimit(limit)
|
|
if len(flattened) != 1 || flattened[0].WindowDurationMinutes != 10080 {
|
|
t.Fatalf("flattened limit = %#v; want a 10080-minute window", flattened)
|
|
}
|
|
}
|
|
|
|
type fakeCodexClient struct {
|
|
mu sync.Mutex
|
|
connected bool
|
|
starts int
|
|
initializes int
|
|
closes int
|
|
calls int
|
|
initErrors []error
|
|
initStarted chan struct{}
|
|
initRelease chan struct{}
|
|
}
|
|
|
|
func (f *fakeCodexClient) Start(context.Context) error {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.starts++
|
|
f.connected = true
|
|
return nil
|
|
}
|
|
|
|
func (f *fakeCodexClient) Initialize(context.Context) error {
|
|
f.mu.Lock()
|
|
f.initializes++
|
|
var err error
|
|
if len(f.initErrors) > 0 {
|
|
err, f.initErrors = f.initErrors[0], f.initErrors[1:]
|
|
}
|
|
started, release := f.initStarted, f.initRelease
|
|
f.mu.Unlock()
|
|
if started != nil {
|
|
select {
|
|
case started <- struct{}{}:
|
|
default:
|
|
}
|
|
}
|
|
if release != nil {
|
|
<-release
|
|
}
|
|
return err
|
|
}
|
|
|
|
func (f *fakeCodexClient) Call(_ context.Context, method string, _ any, out any) error {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.calls++
|
|
if method == "account/login/start" {
|
|
result := out.(*map[string]any)
|
|
*result = map[string]any{"verificationUrl": "https://example.test/device", "userCode": "ABCD-EFGH"}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (f *fakeCodexClient) Close() error {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.closes++
|
|
f.connected = false
|
|
return nil
|
|
}
|
|
|
|
func (f *fakeCodexClient) Connected() bool {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
return f.connected
|
|
}
|
|
|
|
func (f *fakeCodexClient) counts() (starts, initializes, closes, calls int) {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
return f.starts, f.initializes, f.closes, f.calls
|
|
}
|
|
|
|
func TestEnsureReadySerializesColdStart(t *testing.T) {
|
|
client := &fakeCodexClient{}
|
|
rt := &accountRuntime{client: client}
|
|
var wg sync.WaitGroup
|
|
errs := make(chan error, 8)
|
|
for range 8 {
|
|
wg.Add(1)
|
|
go func() {
|
|
defer wg.Done()
|
|
errs <- rt.ensureReady(context.Background())
|
|
}()
|
|
}
|
|
wg.Wait()
|
|
close(errs)
|
|
for err := range errs {
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
starts, initializes, _, _ := client.counts()
|
|
if starts != 1 || initializes != 1 {
|
|
t.Fatalf("cold starts = %d, initializes = %d; want 1 each", starts, initializes)
|
|
}
|
|
}
|
|
|
|
func TestEnsureReadyRetriesAfterInitializeFailure(t *testing.T) {
|
|
client := &fakeCodexClient{initErrors: []error{errors.New("handshake failed")}}
|
|
rt := &accountRuntime{client: client}
|
|
if err := rt.ensureReady(context.Background()); err == nil {
|
|
t.Fatal("first initialization unexpectedly succeeded")
|
|
}
|
|
if err := rt.ensureReady(context.Background()); err != nil {
|
|
t.Fatalf("retry failed: %v", err)
|
|
}
|
|
starts, initializes, closes, _ := client.counts()
|
|
if starts != 2 || initializes != 2 || closes != 1 {
|
|
t.Fatalf("starts = %d, initializes = %d, closes = %d; want 2, 2, 1", starts, initializes, closes)
|
|
}
|
|
}
|
|
|
|
func TestStopWaitsForStartupAndPreventsRestart(t *testing.T) {
|
|
started := make(chan struct{}, 1)
|
|
release := make(chan struct{})
|
|
client := &fakeCodexClient{initStarted: started, initRelease: release}
|
|
rt := &accountRuntime{client: client}
|
|
readyDone := make(chan error, 1)
|
|
go func() { readyDone <- rt.ensureReady(context.Background()) }()
|
|
<-started
|
|
stopDone := make(chan struct{})
|
|
go func() { rt.stop(); close(stopDone) }()
|
|
close(release)
|
|
if err := <-readyDone; err != nil {
|
|
t.Fatalf("startup failed: %v", err)
|
|
}
|
|
<-stopDone
|
|
if err := rt.ensureReady(context.Background()); !errors.Is(err, errRuntimeStopped) {
|
|
t.Fatalf("restart error = %v; want stopped", err)
|
|
}
|
|
starts, initializes, closes, _ := client.counts()
|
|
if starts != 1 || initializes != 1 || closes != 1 {
|
|
t.Fatalf("starts = %d, initializes = %d, closes = %d; want 1 each", starts, initializes, closes)
|
|
}
|
|
}
|
|
|
|
func TestDeviceLoginStartsColdRuntime(t *testing.T) {
|
|
client := &fakeCodexClient{}
|
|
a := &App{runtimes: map[int64]*accountRuntime{2: {client: client}}}
|
|
recorder := httptest.NewRecorder()
|
|
request := httptest.NewRequest("POST", "/api/v1/accounts/2/login/device", nil)
|
|
a.deviceLogin(recorder, request, 2)
|
|
if recorder.Code != 200 {
|
|
t.Fatalf("status = %d, body = %s", recorder.Code, recorder.Body.String())
|
|
}
|
|
starts, initializes, _, calls := client.counts()
|
|
if starts != 1 || initializes != 1 || calls != 1 {
|
|
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)
|
|
}
|
|
})
|
|
}
|
|
}
|