feat: validate Team workspace connections

This commit is contained in:
zhoujun0601
2026-08-13 09:03:47 -04:00
parent 6adbb2d714
commit eb7f9fd66c
10 changed files with 345 additions and 28 deletions
+45 -4
View File
@@ -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 {
+36 -2
View File
@@ -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 {
+23
View File
@@ -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)
}
})
}
}