Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
93 changes: 61 additions & 32 deletions cmd/login/main.go
Original file line number Diff line number Diff line change
@@ -1,19 +1,24 @@
// login.go — WorkBuddy CN OAuth 登录(设备授权流程,CN realm only)。
// login.go — WorkBuddy CN + 国际版 OAuth 登录(设备授权流程, realm)。
//
// 两个子命令,由 login.sh 顺序驱动
// 两个子命令,由 login.sh 顺序驱动(加 --realm=cn|global,默认 cn):
//
// login url → POST /v2/plugin/auth/state?platform=CLI 拿 state+authUrl,
// state 落 /tmp/wb2api-login-state.json,stdout 打印授权 URL
// login poll → 读 state,GET /v2/plugin/auth/token?state= 一次,
// 成功再 GET /v2/plugin/login/account?state= 拿 uid/nickname,
// stdout 打印完整 token+account JSON
// login [--realm=global] url → POST {base}/v2/plugin/auth/state?platform=CLI 拿 state+authUrl,
// state 落 /tmp/wb2api-login-state.json,stdout 打印授权 URL
// login [--realm=global] poll → 读 state,GET {base}/v2/plugin/auth/token?state= 一次,
// 成功再 GET {base}/v2/plugin/login/account?state= 拿 uid/nickname,
// stdout 打印完整 token+account JSON(含 realm 字段)
//
// CN:base=https://copilot.tencent.com,Origin=https://www.codebuddy.cn
// 国际版:base=https://www.workbuddy.ai,Origin=https://www.workbuddy.ai
// (实测国际版 auth/state + auth/token 与 CN 同路径同信封)
//
// 无 PKCE(workbuddy 设备流由服务端签发 state)。
package main

import (
"bytes"
"encoding/json"
"flag"
"fmt"
"io"
"net/http"
Expand All @@ -22,25 +27,34 @@ import (
"time"
)

// 上游常量(CN only
// 上游常量(双 realm:CN + 国际版
const (
upstreamBaseCN = "https://copilot.tencent.com"
clientUA = "CLI/2.63.2 CodeBuddy/2.63.2"
originReferer = "https://www.codebuddy.cn"
endpointAuthState = upstreamBaseCN + "/v2/plugin/auth/state?platform=CLI"
endpointLoginAcct = upstreamBaseCN + "/v2/plugin/login/account?state="
endpointAuthToken = upstreamBaseCN + "/v2/plugin/auth/token?state="
stateFile = "/tmp/wb2api-login-state.json"
upstreamBaseCN = "https://copilot.tencent.com"
upstreamBaseGlobal = "https://www.workbuddy.ai"
clientUA = "CLI/2.63.2 CodeBuddy/2.63.2"
originRefererCN = "https://www.codebuddy.cn"
originRefererGlobal = "https://www.workbuddy.ai"
stateFile = "/tmp/wb2api-login-state.json"
)

// commonHeaders 通用请求头
func commonHeaders(req *http.Request) {
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json, text/plain, */*")
req.Header.Set("X-Requested-With", "XMLHttpRequest")
req.Header.Set("Origin", originReferer)
req.Header.Set("Referer", originReferer+"/")
req.Header.Set("User-Agent", clientUA)
// realmConfig 按 realm 返回 base + origin。
func realmConfig(realm string) (base, origin string) {
if realm == "global" {
return upstreamBaseGlobal, originRefererGlobal
}
return upstreamBaseCN, originRefererCN
}

// commonHeaders 通用请求头(origin 按 realm 切换)。
func commonHeaders(origin string) func(*http.Request) {
return func(req *http.Request) {
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json, text/plain, */*")
req.Header.Set("X-Requested-With", "XMLHttpRequest")
req.Header.Set("Origin", origin)
req.Header.Set("Referer", origin+"/")
req.Header.Set("User-Agent", clientUA)
}
}

// apiEnvelope 与 main.go:429-433 一致
Expand All @@ -59,7 +73,7 @@ func doJSON(client *http.Client, method, fullURL string, headers func(*http.Requ
if headers != nil {
headers(req)
} else {
commonHeaders(req)
commonHeaders(originRefererCN)(req)
}
resp, err := client.Do(req)
if err != nil {
Expand Down Expand Up @@ -90,20 +104,30 @@ func fatal(format string, args ...any) {

type loginState struct {
State string `json:"state"`
Realm string `json:"realm,omitempty"`
}

func main() {
if len(os.Args) < 2 {
fatal("usage: login <url|poll>")
fs := flag.NewFlagSet("login", flag.ContinueOnError)
realm := fs.String("realm", "cn", "login realm: cn|global")
_ = fs.Parse(os.Args[1:])
rest := fs.Args()
if len(rest) < 1 {
fatal("usage: login [--realm=cn|global] <url|poll>")
}
base, origin := realmConfig(*realm)
headers := commonHeaders(origin)
endpointAuthState := base + "/v2/plugin/auth/state?platform=CLI"
endpointLoginAcct := base + "/v2/plugin/login/account?state="
endpointAuthToken := base + "/v2/plugin/auth/token?state="
// 每个流程独立 cookie jar(oauth.go:22-29:多账号登录互不串会话)
jar, _ := cookiejar.New(nil)
client := &http.Client{Timeout: 30 * time.Second, Jar: jar}

switch os.Args[1] {
switch rest[0] {
case "url":
// handleStartLogin (oauth.go:68-87)
data, _, err := doJSON(client, http.MethodPost, endpointAuthState, nil, bytes.NewReader([]byte("{}")))
data, _, err := doJSON(client, http.MethodPost, endpointAuthState, headers, bytes.NewReader([]byte("{}")))
if err != nil {
fatal("auth state failed: %v", err)
}
Expand All @@ -114,7 +138,7 @@ func main() {
if err := json.Unmarshal(data, &st); err != nil || st.State == "" || st.AuthURL == "" {
fatal("auth state: missing state or authUrl")
}
raw, _ := json.Marshal(loginState{State: st.State})
raw, _ := json.Marshal(loginState{State: st.State, Realm: *realm})
if err := os.WriteFile(stateFile, raw, 0o600); err != nil {
fatal("write state: %v", err)
}
Expand All @@ -129,9 +153,13 @@ func main() {
if err := json.Unmarshal(raw, &ls); err != nil {
fatal("parse state: %v", err)
}
// url 与 poll 必须同 realm:poll 用 state 落盘时的 realm(防混域)。
if ls.Realm != "" && ls.Realm != *realm {
fatal("realm mismatch: url 用 --realm=%s,poll 也要用 --realm=%s", ls.Realm, ls.Realm)
}
// handlePollLogin (oauth.go:108-162):auth/token 是权威登录状态端点,
// pending 时业务 code 非 0("login ing"),完成时 code=0 + token bundle
tokRaw, status, errTok := doJSON(client, http.MethodGet, endpointAuthToken+ls.State, nil, nil)
tokRaw, status, errTok := doJSON(client, http.MethodGet, endpointAuthToken+ls.State, headers, nil)
if errTok != nil {
if status == 0 || status >= 500 {
fatal("token endpoint error: %v", errTok)
Expand All @@ -154,7 +182,7 @@ func main() {
Nickname string `json:"nickname"`
}
acctHeaders := func(r *http.Request) {
commonHeaders(r)
headers(r)
r.Header.Set("Authorization", "Bearer "+tok.AccessToken)
}
if acctRaw, _, errAcct := doJSON(client, http.MethodGet, endpointLoginAcct+ls.State, acctHeaders, nil); errAcct == nil {
Expand All @@ -165,6 +193,7 @@ func main() {
"refresh_token": tok.RefreshToken,
"expires_in": tok.ExpiresIn,
"domain": tok.Domain,
"realm": *realm,
"uid": acct.UID,
"enterprise_id": acct.EnterpriseID,
"nickname": acct.Nickname,
Expand All @@ -174,6 +203,6 @@ func main() {
os.Remove(stateFile)

default:
fatal("unknown subcommand %q (want url|poll)", os.Args[1])
fatal("unknown subcommand %q (want url|poll)", rest[0])
}
}
14 changes: 14 additions & 0 deletions internal/auth/auth.go
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,20 @@ func (a *Auth) NeedsRefresh(within time.Duration) bool {
return time.Now().Add(within).Unix() >= a.ExpiresAt
}

// IsGlobal 报告该账号是否属于国际版(workbuddy.ai 域)。
// 判定规则:domain 以 .workbuddy.ai 结尾;空 domain 回落 CN(历史 CN 凭证无 domain 字段)。
func (a *Auth) IsGlobal() bool {
return strings.HasSuffix(a.Domain, ".workbuddy.ai")
}

// Realm 返回账号所属域:"global" 或 "cn",供日志/调度区分。
func (a *Auth) Realm() string {
if a.IsGlobal() {
return "global"
}
return "cn"
}

// Parse 兼容两种磁盘形态:
//
// 嵌套形 {"auth":{...},"account":{...}} (插件 OAuth 输出)
Expand Down
27 changes: 27 additions & 0 deletions internal/auth/realm_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
package auth

import "testing"

// IsGlobal 按 domain 后缀判域;空 domain 回落 CN(历史 CN 凭证无该字段)。
func TestIsGlobal(t *testing.T) {
cases := []struct {
domain string
global bool
}{
{"copilot.tencent.com", false}, // CN 登录返回的 domain
{"", false}, // 空 domain 回落 CN
{"www.codebuddy.cn", false},
{"abc.workbuddy.ai", true}, // 国际版
{"www.workbuddy.ai", true},
{"copilot.tencent.com.evil.workbuddy.ai", true}, // 后缀匹配
}
for _, c := range cases {
a := &Auth{Domain: c.domain}
if got := a.IsGlobal(); got != c.global {
t.Errorf("domain=%q: IsGlobal=%v, want %v", c.domain, got, c.global)
}
if want := map[bool]string{true: "global", false: "cn"}[c.global]; a.Realm() != want {
t.Errorf("domain=%q: Realm=%q, want %q", c.domain, a.Realm(), want)
}
}
}
27 changes: 20 additions & 7 deletions login.sh
Original file line number Diff line number Diff line change
@@ -1,19 +1,27 @@
#!/usr/bin/env bash
# login.sh — WorkBuddy CN OAuth 登录 → 落盘 auth 文件
# login.sh — WorkBuddy CN/国际版 OAuth 登录 → 落盘 auth 文件
#
# 用法:
# ./login.sh
# ./login.sh [--realm=cn|global] # 默认 cn;国际版用 --realm=global
#
# 流程:
# 1. POST /v2/plugin/auth/state 拿授权 URL(无 PKCE,state 由服务端签发)
# 1. POST {base}/v2/plugin/auth/state 拿授权 URL(无 PKCE,state 由服务端签发)
# 2. 你在浏览器打开 URL 完成登录
# 3. 回到这里按 y → poll 拿 token+uid+nickname → 签到 → 落盘 auths/workbuddy-<uid>.json
# 3. 回到这里按 y → poll 拿 token+uid+nickname → 签到(仅 CN,国际版跳过)→ 落盘 auths/workbuddy-<uid>.json
# 4. 重启 workbuddy2api 容器加载新账号
set -euo pipefail

cd "$(dirname "$0")"
AUTH_DIR="./auths"
CONTAINER="workbuddy2api"
# realm 参数透传给 login 二进制(默认 cn)
REALM_ARGS=()
REALM="cn"
for arg in "$@"; do
case "$arg" in
--realm=*) REALM="${arg#--realm=}"; REALM_ARGS=("$arg");;
esac
done

mkdir -p "$AUTH_DIR"

Expand All @@ -28,7 +36,7 @@ echo " WorkBuddy OAuth 登录"
echo "============================================================"
echo ""

AUTH_URL=$("$LOGIN_BIN" url)
AUTH_URL=$("$LOGIN_BIN" "${REALM_ARGS[@]}" url)

echo "请在浏览器中打开以下链接完成登录:"
echo ""
Expand All @@ -51,7 +59,7 @@ fi
echo ""
echo "正在获取 token..."

RESULT=$("$LOGIN_BIN" poll) || {
RESULT=$("$LOGIN_BIN" "${REALM_ARGS[@]}" poll) || {
echo ""
echo "获取 token 失败。可能原因:"
echo " - 登录还没完成就按了 y(重新运行 ./login.sh 再试)"
Expand All @@ -74,7 +82,11 @@ fi

EXPIRES_AT=$(( $(date +%s) + EXPIRES_IN ))

# ─── 签到(CN:POST codebuddy.cn/v2/billing/meter/daily-checkin,幂等不阻塞)───
# ─── 签到(仅 CN:POST codebuddy.cn/v2/billing/meter/daily-checkin,幂等不阻塞;
# 国际版无此体系,跳过)───
if [[ "$REALM" == "global" ]]; then
echo "签到: 跳过(国际版无 CN 签到体系)"
else
python3 - <<PYEOF
import json, urllib.request, urllib.error

Expand Down Expand Up @@ -107,6 +119,7 @@ except urllib.error.HTTPError as e:
except Exception as e:
print(f"签到: {e}")
PYEOF
fi

# ─── 落盘 auth 文件(与 internal/auth 读取格式一致)─────────────────
AUTH_FILE="$AUTH_DIR/workbuddy-${USER_ID}.json"
Expand Down