Repository navigation
Expand file tree
/
Copy pathconverter.py
More file actions
1812 lines (1579 loc) · 75.8 KB
/
Copy pathconverter.py
File metadata and controls
1812 lines (1579 loc) · 75.8 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
#!/usr/bin/env python3
"""
workbuddy2api — 把 CodeBuddy / WorkBuddy 的订阅暴露成标准 OpenAI 兼容 API。
原理(直连后端,原生 function calling):
- 读取本机已登录的 CodeBuddy 桌面端凭据(auth 文件里的 token / uid / enterpriseId)。
- 直接转发到 CodeBuddy 后端 `https://copilot.tencent.com/v2/chat/completions`。
该后端本身就是标准 OpenAI chat/completions 协议(含原生 tools / tool_calls / SSE 流式)。
- 转换器只做两件事:①注入鉴权 header(Authorization / X-User-Id 等)
②在本地 /v1/* 与后端 /v2/* 之间做路径映射与透传(含 Anthropic / Chat / Responses 三种协议)。
- token 过期时自动调 `/v2/plugin/auth/token/refresh` 刷新,并回写 auth 文件。
跨平台:自动定位 auth 目录(macOS / Windows / Linux)。
依赖:fastapi + uvicorn + httpx(pip install fastapi "uvicorn[standard]" httpx)。
用法:
python3 converter.py # 默认 127.0.0.1:8787
python3 converter.py --port 9000
python3 converter.py --api-key mysecret # 启用客户端鉴权
"""
from __future__ import annotations
import argparse
import base64
import hashlib
import json
import os
import re
import sys
import threading
import time
from pathlib import Path
from typing import Optional
import httpx
from fastapi import FastAPI, Header, HTTPException, Request
from fastapi.responses import JSONResponse, StreamingResponse
import uvicorn
# 连接池:减少重复 TLS 握手; MaxIdleConnsPerHost=20 设计。
_HTTP_LIMITS = httpx.Limits(max_connections=100, max_keepalive_connections=20)
def _client_ip_headers(request: Request, purpose: str = "conversation") -> dict:
"""提取真实客户端 IP 与用途/产品头,透传给上游,避免请求用量里 client/agentPurpose 为空。
真实 WorkBuddy 桌面端:
- X-Agent-Purpose: "conversation" 用于普通对话
- X-IDE-Name / X-IDE-Type / X-Product: "WorkBuddy" 用于上游识别 client
"""
ip = None
xff = request.headers.get("X-Forwarded-For")
if xff:
ip = xff.split(",")[0].strip()
else:
real = request.headers.get("X-Real-IP")
if real:
ip = real.strip()
elif request.client:
ip = request.client.host
client_name = os.environ.get("ADMIN_UPSTREAM_CLIENT_NAME", "WorkBuddy").strip() or "WorkBuddy"
h = {
"X-Agent-Purpose": purpose or "conversation",
"X-IDE-Name": client_name,
"X-IDE-Type": client_name,
"X-Product": client_name,
}
if ip:
h["X-Forwarded-For"] = ip
h["X-Real-IP"] = ip
h["X-Client-IP"] = ip
custom = os.environ.get("ADMIN_UPSTREAM_CLIENT_HEADER", "").strip()
if custom:
h[custom] = ip
return h
try:
from desensitize import desensitize_body
except ImportError: # 模块缺失时降级为不脱敏
def desensitize_body(body, roles=("system",), desensitize_harness_user=False,
desensitize_tools=False, compact_harness=False,
strip_tool_metadata=False):
return body
from responses_adapter import (
responses_request_to_chat,
ResponsesStreamConverter,
)
from responses_projection import project_responses_chat_body
from anthropic_adapter import (
anthropic_request_to_chat,
AnthropicStreamConverter,
)
# ---------------------------------------------------------------------------
# 常量
# ---------------------------------------------------------------------------
BACKEND = "https://copilot.tencent.com"
DEFAULT_DOMAIN = "www.codebuddy.cn"
USER_AGENT = "codebuddy2openai/2.0"
# ---------------------------------------------------------------------------
# 平台相关:定位 auth 目录
# ---------------------------------------------------------------------------
def auth_dirs() -> list[Path]:
env_dir = os.environ.get("CODEBUDDY_AUTH_DIR")
if env_dir:
return [Path(env_dir)]
home = Path.home()
plat = sys.platform
if plat == "darwin":
return [home / "Library" / "Application Support" / "CodeBuddyExtension" / "Data" / "Public" / "auth"]
if plat == "win32":
local = Path(os.environ.get("LOCALAPPDATA", home / "AppData" / "Local"))
return [local / "CodeBuddyExtension" / "Data" / "Public" / "auth"]
xdg = Path(os.environ.get("XDG_DATA_HOME", home / ".local" / "share"))
return [xdg / "CodeBuddyExtension" / "Data" / "Public" / "auth"]
def find_auth_file() -> Path | None:
for d in auth_dirs():
if not d.is_dir():
continue
# 优先用桌面端实时登录文件(无时间戳后缀),避免被历史备份按字典序抢走
live = d / "workbuddy-desktop.info"
if live.is_file():
return live
files = sorted(d.glob("*.info"))
if files:
return files[0]
return None
# ---------------------------------------------------------------------------
# Auth 凭据管理(读 + 自动刷新 + 回写)
# ---------------------------------------------------------------------------
def _get_turing_device_token() -> str | None:
"""延迟取本机设备风控 Token;失败返回 None(不影响主流程)。
放在模块级做懒加载:converter.py 既可作为 admin 的子模块被挂载,也可独立
`python converter.py` 运行。admin 包不可用时(极少数情况)直接降级为不带该头。
"""
try:
from admin.turing_token import get_device_token
return get_device_token()
except Exception as error:
_log(f'获取device token错误: {error}')
return None
def _maybe_decrypt_auth_file(raw: dict) -> dict | None:
"""递归解密 auth JSON 中所有 $wbEncrypted:1 字段。若有变更返回新 dict,否则 None。"""
try:
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
except ImportError as e:
raise RuntimeError("解密 auth 文件需要 cryptography 包,请安装 requirements.txt") from e
_at_rest_secret_key = "Sik9U5aXhCdwTVEwsEySDOmDoB9r9ntFxHF1fst9LQI="
_key = hashlib.sha256(_at_rest_secret_key.encode("utf-8")).digest()
_key_id = hashlib.sha256(_key).hexdigest()[:16]
def _lp(s: str) -> bytes:
b = s.encode("utf-8")
return len(b).to_bytes(4, "big") + b
def _build_aad(suite: int, key_id: str) -> bytes:
return (
b"WB-AAD\0"
+ b"\x01"
+ _lp("WBEV1")
+ _lp("sym-v1")
+ suite.to_bytes(4, "big")
+ _lp(key_id)
+ b"\x02"
+ b"\x00"
+ b"\x00"
)
def _unwrap(value):
if (
isinstance(value, dict)
and value.get("type") == "Buffer"
and isinstance(value.get("data"), list)
):
return bytes(value["data"]).decode("utf-8")
if isinstance(value, dict):
return {k: _unwrap(v) for k, v in value.items()}
if isinstance(value, list):
return [_unwrap(v) for v in value]
return value
def _maybe_json(s: str):
try:
return json.loads(s)
except json.JSONDecodeError:
return s
def _decrypt_field(envelope_b64: str):
env = json.loads(base64.b64decode(envelope_b64).decode("utf-8"))
if env.get("keyId") != _key_id:
raise RuntimeError(f"auth 文件 envelope keyId 不匹配:{env.get('keyId')}")
nonce = base64.b64decode(env["nonce"])
tag = base64.b64decode(env["authTag"])
ct = base64.b64decode(env["ciphertext"])
aad = _build_aad(env["suite"], env["keyId"])
plaintext = AESGCM(_key).decrypt(nonce, ct + tag, aad)
v = _maybe_json(plaintext.decode("utf-8"))
v = _unwrap(v)
if isinstance(v, str):
return _maybe_json(v)
return v
def _walk(value):
if isinstance(value, dict):
if value.get("$wbEncrypted") == 1 and isinstance(value.get("envelope"), str):
return (_decrypt_field(value["envelope"]), True)
new_dict: dict = {}
changed = False
for k, v in value.items():
nv, c = _walk(v)
new_dict[k] = nv
changed = changed or c
return (new_dict, changed)
if isinstance(value, list):
new_list: list = []
changed = False
for v in value:
nv, c = _walk(v)
new_list.append(nv)
changed = changed or c
return (new_list, changed)
return (value, False)
new_raw: dict = {}
changed = False
for k, v in raw.items():
nv, c = _walk(v)
new_raw[k] = nv
changed = changed or c
return new_raw if changed else None
class CredentialManager:
"""从 auth 文件读取凭据;token 临近过期时自动刷新并回写。"""
def __init__(self, path: Path):
self.path = path
self._lock = threading.Lock()
self._cached: dict | None = None
self._mtime: float = 0.0
try:
self._cached = self._read_decrypted()
self._mtime = self.path.stat().st_mtime
except (OSError, json.JSONDecodeError) as e:
_log(f"读取 auth 文件失败:{e}")
def _read_raw(self) -> dict:
with open(self.path, "r", encoding="utf-8") as f:
return json.load(f)
def _read_decrypted(self) -> dict:
"""读取 auth 文件;若含加密字段则解密后返回,不修改原文件。"""
raw = self._read_raw()
decrypted = _maybe_decrypt_auth_file(raw)
if decrypted is None:
return raw
_log("解密 auth 文件中的加密字段")
return decrypted
def _load_if_stale(self):
"""若文件 mtime 变了(外部刷新过),重新加载缓存。"""
try:
mt = self.path.stat().st_mtime
except OSError:
return
if self._cached is None or mt != self._mtime:
try:
self._cached = self._read_decrypted()
self._mtime = self.path.stat().st_mtime
except (OSError, json.JSONDecodeError) as e:
_log(f"读取 auth 文件失败:{e}")
def _session(self) -> dict:
self._load_if_stale()
if self._cached is None:
raise RuntimeError(f"无法读取 auth 文件:{self.path}")
return self._cached
def _is_expired(self) -> bool:
s = self._session()
expires_at = (s.get("auth") or {}).get("expiresAt") or 0
# 提前 60s 判定过期
return time.time() * 1000 >= (expires_at - 60_000)
def _refresh(self):
"""调后端刷新 token,写回 auth 文件与缓存。"""
s = self._session()
auth = s.get("auth") or {}
headers = self._build_headers_from(auth, s.get("account") or {})
headers["X-Refresh-Token"] = auth.get("refreshToken", "")
headers["X-Auth-Refresh-Source"] = "plugin"
url = f"{BACKEND}/v2/plugin/auth/token/refresh"
try:
with httpx.Client(timeout=15, limits=_HTTP_LIMITS) as c:
r = c.post(url, headers=headers, json={})
data = r.json()
except Exception as e:
raise RuntimeError(f"刷新 token 网络失败:{e}")
if data.get("code") != 0 or not data.get("data"):
raise RuntimeError(f"刷新 token 失败:{data.get('msg', data)}")
new_auth = data["data"]
# 继承部分字段
new_auth["domain"] = new_auth.get("domain") or auth.get("domain")
new_auth["lastRefreshTime"] = int(time.time() * 1000)
# 计算 expiresAt(若后端没直接给)
if not new_auth.get("expiresAt") and new_auth.get("expiresIn"):
new_auth["expiresAt"] = int(time.time() * 1000) + new_auth["expiresIn"] * 1000
if not new_auth.get("refreshExpiresAt") and new_auth.get("refreshExpiresIn"):
new_auth["refreshExpiresAt"] = int(time.time() * 1000) + new_auth["refreshExpiresIn"] * 1000
s["auth"] = new_auth
# 原子写回
tmp = self.path.with_suffix(self.path.suffix + ".tmp")
with open(tmp, "w", encoding="utf-8") as f:
json.dump(s, f, ensure_ascii=False, indent=2)
os.replace(tmp, self.path)
self._cached = s
self._mtime = self.path.stat().st_mtime
def _build_headers_from(self, auth: dict, account: dict) -> dict:
domain = auth.get("domain") or DEFAULT_DOMAIN
h = {
"Content-Type": "application/json",
"Accept": "application/json",
"Authorization": f"Bearer {auth.get('accessToken','')}",
"X-User-Id": account.get("uid", ""),
"X-Enterprise-Id": account.get("enterpriseId", ""),
"X-Tenant-Id": account.get("enterpriseId", ""),
"X-Domain": domain,
"User-Agent": USER_AGENT,
}
# 风控设备头:与桌面端 Turing Shield 一致,缺失会被上游识别为异常客户端。
# 取不到(桌面端未安装 / SDK 不支持)时优雅降级为不带该头,不影响主流程。
tok = _get_turing_device_token()
if tok:
_log(f'获取到device token:{tok}')
h["X-Device-Token"] = tok
return h
def get_headers(self, extra: dict | None = None) -> dict:
"""返回带最新 token 的后端请求 header;必要时先刷新。
extra: 调用方(如 proxy.py)可注入的风控/审计头,例如真实客户端 IP、
用途标识 X-Agent-Purpose 等。这些头会被 merge 到基础鉴权头之后,
确保上游请求用量能正确显示 client 与 agentPurpose,降低被风控概率。
"""
with self._lock:
if self._is_expired():
self._refresh()
s = self._session()
h = self._build_headers_from(s.get("auth") or {}, s.get("account") or {})
if extra:
h.update(extra)
return h
def summary(self) -> dict:
s = self._session()
auth = s.get("auth") or {}
acct = s.get("account") or {}
exp = auth.get("expiresAt", 0)
return {
"uid": acct.get("uid"),
"nickname": acct.get("nickname"),
"enterpriseName": acct.get("enterpriseName"),
"token_expires_at": exp,
"token_expired": self._is_expired(),
}
# -----------------------------------------------------------------------
# 后端资源查询(模型列表、额度)
# -----------------------------------------------------------------------
def _request_backend(self, method: str, path: str, json_body: dict | None = None) -> dict:
"""向后端发一个同步请求,返回 {code, msg, requestId, data} 或抛异常。"""
headers = self.get_headers()
url = f"{BACKEND}{path}"
try:
with httpx.Client(timeout=15, limits=_HTTP_LIMITS) as c:
if method.upper() == "GET":
r = c.get(url, headers=headers)
else:
r = c.post(url, headers=headers, json=json_body or {})
except Exception as e:
raise RuntimeError(f"后端请求网络失败 {method} {path}: {e}")
try:
data = r.json()
except Exception as e:
raise RuntimeError(f"后端返回非 JSON {method} {path} HTTP {r.status_code}: {r.text[:200]}")
if r.status_code != 200 or data.get("code") != 0:
raise RuntimeError(f"后端请求失败 {method} {path}: HTTP {r.status_code} / {data.get('msg', data)}")
return data
def _request_backend_soft(self, method: str, path: str, json_body: dict | None = None) -> dict:
"""同 `_request_backend`,但后端返回业务 code!=0 时**不抛异常**,原样返回解析后的 dict。
用于签到领取等场景:领取接口的 1001(已领)/1002(无资格)/1003(活动结束) 等业务码
属于正常业务结果,需要由调用方根据 code 区分处理,而非当作错误抛掉。
"""
headers = self.get_headers()
url = f"{BACKEND}{path}"
try:
with httpx.Client(timeout=15, limits=_HTTP_LIMITS) as c:
if method.upper() == "GET":
r = c.get(url, headers=headers)
else:
r = c.post(url, headers=headers, json=json_body or {})
except Exception as e:
raise RuntimeError(f"后端请求网络失败 {method} {path}: {e}")
try:
return r.json()
except Exception as e:
raise RuntimeError(f"后端返回非 JSON {method} {path} HTTP {r.status_code}: {r.text[:200]}")
def _enterprise_path_key(self) -> str:
"""返回模型列表 endpoint 里的 enterprise 段:personal 或 enterpriseId。"""
s = self._session()
acct = s.get("account") or {}
if acct.get("type") == "personal":
return "personal"
eid = acct.get("enterpriseId")
return eid if eid else "personal"
@staticmethod
def _parse_credit_multiplier(credits) -> float | None:
"""把 'x0.05' / 'x0.00 credits' 解析成浮点倍率,解析不出返回 None。"""
if not credits:
return None
m = re.search(r"x\s*([0-9]+(?:\.[0-9]+)?)", str(credits))
return float(m.group(1)) if m else None
def fetch_models(self) -> list[dict]:
"""获取后端真实模型列表(含 id/name/credits 等元信息)。
兼容两种返回结构:
- 顶层 data.models 为对象数组(每项含 id/name/credits...);
- 仅 data.agents[].models 为字符串数组时,回退收集并去重。
"""
eid = self._enterprise_path_key()
data = self._request_backend("GET", f"/v2/enterprises/{eid}/models")
payload = data.get("data", {})
models = payload.get("models")
if isinstance(models, list) and models:
return models
# 兜底:从 agents 里收集模型名
collected: list[dict] = []
for a in payload.get("agents", []) or []:
for m in a.get("models", []) or []:
if isinstance(m, dict) and m.get("id"):
collected.append(m)
elif isinstance(m, str) and m:
collected.append({"id": m})
seen = set()
result: list[dict] = []
for m in collected:
mid = m.get("id")
if mid and mid not in seen:
seen.add(mid)
result.append(m)
if not result:
raise RuntimeError("后端模型列表格式异常:缺少 data.models 且 agents 中无模型")
return result
def fetch_balance(self) -> dict:
"""获取当前账号积分汇总,仅返回总量与剩余(可用积分)。"""
data = self._request_backend("POST", "/v2/billing/meter/get-user-resource", {})
resp = data.get("data", {}).get("Response", {}).get("Data", {}) or {}
total = 0
total_size = 0
for a in resp.get("Accounts") or []:
if a.get("CapacityUnit") != "credits":
continue
total += a.get("CapacityRemain") or 0
total_size += a.get("CapacitySize") or 0
return {
"total": total_size, # 总积分
"remain": total, # 可用积分(剩余额度)
}
# ---------------------------------------------------------------------------
# 模型列表
# ---------------------------------------------------------------------------
DEFAULT_MODELS = [
"glm-5.2", "glm-5.1", "glm-5v-turbo",
"kimi-k2.7", "kimi-k2.6", "kimi-k2.5",
"deepseek-v4-pro", "deepseek-v4-flash",
"minimax-m3-pay", "hy3-preview-agent", "auto",
]
# 后端资源缓存(TTL,秒)
_RESOURCE_CACHE_TTL = 60.0
_MODELS_CACHE = {"ts": 0.0, "data": None, "error": None}
_BALANCE_CACHE = {"ts": 0.0, "data": None, "error": None}
# 后端请求体里出现过的额外字段(透传时若客户端给了就保留)
PASSTHROUGH_BODY_KEYS = {
"model", "messages", "tools", "tool_choice", "temperature",
"max_tokens", "max_completion_tokens", "top_p", "stream",
"stream_options", "stop", "presence_penalty", "frequency_penalty",
"n", "response_format", "seed", "user", "reasoning_effort",
"verbosity", "reasoning_summary",
}
# ---------------------------------------------------------------------------
# FastAPI 应用
# ---------------------------------------------------------------------------
app = FastAPI(title="codebuddy2openai", version="2.0")
CONFIG: dict = {"api_key": "", "cred": None, "log_path": None,
"desensitize": False, "no_compact": False,
# 去掉流式 delta 里的空 content:""(GLM 等模型的 reasoning 周期
# 会被 AI SDK 当成"文本开始",提前掐断 Thought 周期,产生上百个
# "Thought for 2ms"。默认开,可用环境变量 WORKBUDDY_STRIP_EMPTY_DELTA=0 关闭)
"strip_empty_delta": os.environ.get("WORKBUDDY_STRIP_EMPTY_DELTA", "1") not in ("0", "false", "no"),
# 把流式 reasoning_content 零散分片在网关层合并成"一整段",统一在首个
# content/tool_calls/finish delta 之前释放,并从 tool_calls 的 arguments
# 流中彻底剔除任何混入的 reasoning 字符。解决两类问题:
# 1) 客户端把每个零散 reasoning 分片渲染成独立 Thought 块 → 几百个
# "Thought for 2ms"(上面 strip_empty_delta 只挡空 content,不够);
# 2) reasoning token 与 tool_call arguments 互相穿插 → 工具参数 JSON
# 被截断/污染(Expected '}' / Unterminated string / Expected 'id'...)。
# 默认开,可用环境变量 WORKBUDDY_COALESCE_REASONING=0 关闭(关闭=原样逐事件转发)。
"coalesce_reasoning": os.environ.get("WORKBUDDY_COALESCE_REASONING", "1") not in ("0", "false", "no")}
# cred: CredentialManager | None
# ---------------------------------------------------------------------------
# 日志(写文件)
# ---------------------------------------------------------------------------
_LOG_LOCK = threading.Lock()
def _log(msg: str):
"""写一行日志到 CONFIG['log_path'] 指定的文件(追加,带时间戳)。未设置则丢弃。"""
path = CONFIG.get("log_path")
if not path:
return
line = f"[{time.strftime('%Y-%m-%d %H:%M:%S')}] {msg}\n"
try:
with _LOG_LOCK:
with open(path, "a", encoding="utf-8") as f:
f.write(line)
except OSError:
pass # 日志失败不应影响主流程
_LOG_OK_LOCK = threading.Lock()
_LOG_BAD_LOCK = threading.Lock()
def _log_ok(msg: str):
line = f"[{time.strftime('%Y-%m-%d %H:%M:%S')}] {msg}\n"
try:
with _LOG_OK_LOCK:
with open('ok.log', "a", encoding="utf-8") as f:
f.write(line)
except OSError:
pass # 日志失败不应影响主流程
def _log_bad(msg: str):
line = f"[{time.strftime('%Y-%m-%d %H:%M:%S')}] {msg}\n"
try:
with _LOG_BAD_LOCK:
with open('bad.log', "a", encoding="utf-8") as f:
f.write(line)
except OSError:
pass # 日志失败不应影响主流程
def _truncate(s: str, n: int = 80) -> str:
s = str(s).replace("\n", " ").strip()
return s[:n] + ("…" if len(s) > n else "")
def _check_auth(authorization: Optional[str], x_api_key: Optional[str]):
key = CONFIG["api_key"]
if not key:
return
token = ""
if authorization and authorization.startswith("Bearer "):
token = authorization[7:].strip()
if not token and x_api_key:
token = x_api_key
if token != key:
raise HTTPException(status_code=401, detail={"error": {"message": "invalid api key", "type": "auth_error"}})
def _cred() -> CredentialManager:
if CONFIG["cred"] is None:
raise HTTPException(status_code=503, detail={"error": {"message": "未找到登录凭据,请先在桌面端登录 CodeBuddy/WorkBuddy", "type": "auth_error"}})
return CONFIG["cred"]
def _cached_models(cred) -> list[dict]:
"""带 TTL 缓存的真实模型列表;失败时抛异常由调用方回退。"""
global _MODELS_CACHE
now = time.time()
if _MODELS_CACHE["data"] is not None and now - _MODELS_CACHE["ts"] < _RESOURCE_CACHE_TTL:
return _MODELS_CACHE["data"]
try:
models = cred.fetch_models()
except Exception as e:
_MODELS_CACHE["error"] = str(e)
raise
_MODELS_CACHE = {"ts": now, "data": models, "error": None}
return models
def _cached_balance(cred) -> dict:
"""带 TTL 缓存的真实积分额度;失败时抛异常由调用方回退。"""
global _BALANCE_CACHE
now = time.time()
if _BALANCE_CACHE["data"] is not None and now - _BALANCE_CACHE["ts"] < _RESOURCE_CACHE_TTL:
return _BALANCE_CACHE["data"]
try:
balance = cred.fetch_balance()
except Exception as e:
_BALANCE_CACHE["error"] = str(e)
raise
_BALANCE_CACHE = {"ts": now, "data": balance, "error": None}
return balance
@app.get("/health")
def health():
cred = CONFIG["cred"]
info: dict = {"status": "ok", "platform": sys.platform, "python": sys.version.split()[0],
"auth_file": str(find_auth_file() or "(未找到)"), "mode": "direct-proxy (native function calling)"}
if cred is not None:
try:
info["credential"] = cred.summary()
except Exception as e:
info["credential_error"] = str(e)
try:
info["balance"] = _cached_balance(cred)
except Exception as e:
info["balance_error"] = str(e)
return info
@app.get("/v1/models")
def list_models(authorization: Optional[str] = Header(default=None),
x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")):
_check_auth(authorization, x_api_key)
cred = CONFIG["cred"]
if cred is not None:
try:
models = _cached_models(cred)
data = [{
"id": m.get("id"),
"object": "model",
"created": 1700000000,
"owned_by": "codebuddy",
"name": m.get("name") or m.get("id"),
"credits": m.get("credits"),
"credit_multiplier": cred._parse_credit_multiplier(m.get("credits")),
"description": m.get("descriptionZh") or m.get("descriptionEn"),
"supports_images": m.get("supportsImages"),
"supports_reasoning": m.get("supportsReasoning"),
"supports_tool_call": m.get("supportsToolCall"),
"max_input_tokens": m.get("maxInputTokens"),
"max_output_tokens": m.get("maxOutputTokens"),
"vendor": m.get("vendor"),
} for m in models if m.get("id")]
return {"object": "list", "data": data, "source": "backend"}
except Exception as e:
_log(f"获取真实模型列表失败,回退到 DEFAULT_MODELS: {e}")
data = [{"id": m, "object": "model", "created": 1700000000, "owned_by": "codebuddy"}
for m in DEFAULT_MODELS]
return {"object": "list", "data": data, "source": "fallback"}
@app.get("/v1/balance")
def get_balance(authorization: Optional[str] = Header(default=None),
x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")):
_check_auth(authorization, x_api_key)
cred = _cred()
try:
return {"object": "balance", "source": "backend", **_cached_balance(cred)}
except Exception as e:
raise HTTPException(status_code=502, detail={"error": {"message": f"获取额度失败:{e}", "type": "upstream_error"}})
@app.post("/v1/chat/completions")
async def chat_completions(request: Request,
authorization: Optional[str] = Header(default=None),
x_api_key: Optional[str] = Header(default=None, alias="X-Api-Key")):
_check_auth(authorization, x_api_key)
cred = _cred()
try:
payload = await request.json()
except Exception as e:
raise HTTPException(status_code=400, detail={"error": {"message": f"bad json: {e}", "type": "invalid_request_error"}})
messages = payload.get("messages") or []
if not messages:
raise HTTPException(status_code=400, detail={"error": {"message": "messages is required", "type": "invalid_request_error"}})
# 构造后端 body:只透传已知的合法字段
client_wants_stream = bool(payload.get("stream"))
body = {k: payload[k] for k in PASSTHROUGH_BODY_KEYS if k in payload}
body.setdefault("model", "auto")
# 后端只支持流式:始终以 stream=True 调后端,非流式由转换器聚合
body["stream"] = True
if "stream_options" not in body:
body["stream_options"] = {"include_usage": True}
# 可选:脱敏。缓解客户端合规模板(如 Codex CLI / ZCode 注入的说明文字)被后端误判为敏感词。
# 处理 system / developer 消息、Codex 注入的上下文 user 消息,以及 tools 的 description。
if CONFIG.get("desensitize"):
body = desensitize_body(body, roles=("system", "assistant"),
desensitize_harness_user=True,
desensitize_tools=True,
compact_harness=not CONFIG.get("no_compact"),
strip_tool_metadata=True)
# 日志:请求摘要
model_name = payload.get("model", "auto")
tool_names = [t.get("function", {}).get("name") for t in (payload.get("tools") or [])
if isinstance(t, dict)]
last_user = _last_user_text(messages)
rid = os.urandom(4).hex()
_log(f"[{rid}] ▶ REQUEST {model_name} | stream={client_wants_stream} | msgs={len(messages)}"
+ (f" | tools={tool_names}" if tool_names else "")
+ (f" | last_user={_truncate(last_user, 60)!r}" if last_user else ""))
# 完整请求体(发往后端的实际内容;若启用脱敏,这里已是脱敏后)
_log(f"[{rid}] ── REQUEST BODY (发往后端) ──\n{json.dumps(body, ensure_ascii=False, indent=2)}")
headers = cred.get_headers()
headers.update(_client_ip_headers(request))
url = f"{BACKEND}/v2/chat/completions"
t0 = time.time()
if client_wants_stream:
return StreamingResponse(
_stream_upstream(url, headers, body, model_name, t0, rid),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
)
# 非流式:后端只支持流式,这里把后端 SSE 聚合成单个 chat.completion 响应
try:
async with httpx.AsyncClient(timeout=300, limits=_HTTP_LIMITS) as c:
async with c.stream("POST", url, headers=headers, json=body) as r:
if r.status_code != 200:
raw = await r.aread()
_log(f"[{rid}] ✗ HTTP {r.status_code} | {model_name} | {_truncate(raw.decode('utf-8','replace'),200)}")
_log(f"[{rid}] ── ERROR BODY ──\n{raw.decode('utf-8','replace')}")
raise HTTPException(status_code=r.status_code, detail=_safe_err_raw(raw, r.status_code))
collected = await _collect_stream(r)
except HTTPException:
raise
except httpx.HTTPError as e:
_log(f"[{rid}] ✗ 网络错误 | {model_name} | {e}")
raise HTTPException(status_code=502, detail={"error": {"message": f"upstream error: {e}", "type": "upstream_error"}})
_log_finish(model_name, t0, collected, rid)
return JSONResponse(content=collected)
def _last_user_text(messages: list) -> str:
"""取最后一条 user 消息的文本,用于日志预览。"""
for m in reversed(messages):
if m.get("role") != "user":
continue
content = m.get("content", "")
if isinstance(content, list):
for blk in content:
if isinstance(blk, dict) and blk.get("type") == "text":
return str(blk.get("text", ""))
return ""
return str(content)
return ""
def _log_finish(model_name: str, t0: float, result: dict, rid: str = ""):
"""记录一次完成的请求:耗时 / finish_reason / usage / 工具调用 / 审核拦截 + 完整响应。"""
elapsed = time.time() - t0
prefix = f"[{rid}] " if rid else ""
choice = (result.get("choices") or [{}])[0]
finish = choice.get("finish_reason")
msg = choice.get("message") or {}
tcs = msg.get("tool_calls") or []
usage = result.get("usage") or {}
tag = ""
if finish == "content-filter":
tag = " ⚠️内容审核拦截"
tc_names = [t.get("function", {}).get("name") for t in tcs]
_log(f"{prefix}◀ RESPONSE {model_name} | {elapsed:.1f}s | finish={finish}{tag}"
+ (f" | tool_calls={tc_names}" if tc_names else "")
+ f" | tokens={usage.get('total_tokens', '?')}")
# 完整响应体
_log(f"{prefix}── RESPONSE BODY ──\n{json.dumps(result, ensure_ascii=False, indent=2)}")
async def _collect_stream(response: httpx.Response) -> dict:
"""消费后端的 OpenAI SSE 流,聚合成单个非流式 chat.completion 对象。
合并所有 chunk 的 delta(content / tool_calls),并取 usage / finish_reason。
"""
content_parts: list[str] = []
# tool_calls: index -> {id, name, arguments(分片拼接)}
tool_calls: dict[int, dict] = {}
model: str | None = None
finish_reason: str | None = None
usage: dict | None = None
async for line in response.aiter_lines():
line = line.strip()
if not line or not line.startswith("data:"):
continue
data = line[5:].strip()
if data == "[DONE]":
break
try:
chunk = json.loads(data)
except json.JSONDecodeError:
continue
model = chunk.get("model") or model
if chunk.get("usage"):
usage = chunk["usage"]
for choice in chunk.get("choices") or []:
if choice.get("finish_reason"):
finish_reason = choice["finish_reason"]
delta = choice.get("delta") or {}
if delta.get("content"):
content_parts.append(delta["content"])
for tc in delta.get("tool_calls") or []:
idx = tc.get("index", 0)
slot = tool_calls.setdefault(idx, {"id": None, "name": None, "arguments": ""})
if tc.get("id"):
slot["id"] = tc["id"]
fn = tc.get("function") or {}
if fn.get("name"):
slot["name"] = fn["name"]
if fn.get("arguments"):
slot["arguments"] += fn["arguments"]
tcs = None
if tool_calls:
tcs = [
{"id": v["id"], "type": "function",
"function": {"name": v["name"], "arguments": v["arguments"]}}
for _, v in sorted(tool_calls.items())
]
finish_reason = finish_reason or "tool_calls"
message = {"role": "assistant", "content": "".join(content_parts) or None}
if tcs:
message["tool_calls"] = tcs
return {
"id": "chatcmpl-" + os.urandom(12).hex(),
"object": "chat.completion",
"created": int(time.time()),
"model": model or "unknown",
"choices": [{"index": 0, "message": message,
"finish_reason": finish_reason or "stop"}],
"usage": usage or {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0},
}
def _safe_err_raw(raw: bytes, status: int) -> dict:
try:
return json.loads(raw.decode("utf-8", "replace"))
except Exception:
return {"error": {"message": raw.decode("utf-8", "replace")[:500], "type": "upstream_error", "code": status}}
async def _stream_upstream(url: str, headers: dict, body: dict,
model_name: str = "?", t0: float = 0.0, rid: str = ""):
"""把后端 SSE 原样转发给客户端(后端已是标准 OpenAI SSE,含 tool_calls)。
同时轻量解析流,统计 finish_reason / tool_calls / usage 用于日志,不阻塞转发。
完整原始 SSE 累积后落盘到日志(调试用)。
当 CONFIG['strip_empty_delta'] 开启时,对每个完整 SSE event 做 delta 清洗:
去掉 `delta.content == ""` / `delta.reasoning_content == ""` 这类"空字段",
避免 AI SDK 把空 content 当成"文本已开始"而提前结束 reasoning 周期。
当 CONFIG['coalesce_reasoning'] 开启时,把所有转发事件再过一遍 _ReasoningCoalescer:
把零散 reasoning 分片合并成"一整段"再释放,并从 tool_calls 参数流里剔除混入的
reasoning,避免几百个 "Thought for Xms" 与工具参数 JSON 被截断/污染。
"""
finish_reason = None
tool_names: list[str] = []
usage: dict = {}
saw_filter = False
line_buf = _SseLineBuffer()
raw_parts: list[bytes] = [] # 累积完整原始 SSE
forwarded_parts: list[bytes] = [] # 累积转发给客户端的 SSE(已清洗)
prefix = f"[{rid}] " if rid else ""
strip = bool(CONFIG.get("strip_empty_delta"))
coal = _ReasoningCoalescer()
def _process_event_lines(lines: list[bytes]) -> bytes | None:
"""处理一个完整 SSE event(以 \\n\\n 结尾的若干行)。
返回要转发给客户端的字节;返回 None 表示该 event 应被丢弃。
"""
if not lines:
return None
new_lines: list[bytes] = []
for ln in lines:
stripped = ln.lstrip()
if strip and stripped.startswith(b"data:"):
payload = stripped[5:].lstrip()
try:
obj = json.loads(payload)
except (json.JSONDecodeError, ValueError):
new_lines.append(ln)
continue
_, new_obj = _sanitize_delta_obj(obj)
if obj is not new_obj:
new_lines.append(b"data: " + json.dumps(
new_obj, ensure_ascii=False, separators=(",", ":")
).encode("utf-8"))
else:
new_lines.append(ln)
continue
new_lines.append(ln)
if not new_lines:
return None
return b"\n".join(new_lines) + b"\n\n"
# —— 关键修复:SSE 事件分组状态必须跨 chunk 持久,绝不能每个 aiter_bytes chunk
# 都重建 sink。此前 `out, feed_line = _make_sink()` 放在 chunk 循环内,每次
# chunk 都新建,导致"一个 SSE event 的行没在本 chunk 看到结尾空行就被丢弃"。
# 当后端一个事件跨多个 TCP chunk(GLM/深度推理模型很常见)时,会随机整段丢掉
# data 事件 → tool_call 的 arguments 缺字符/缺字段、reasoning 缺片段,且随网络
# 时序时好时坏。现在 sink 全局只建一次,靠 _drain() 消费后就清空。
def _make_sink() -> tuple[list[bytes], callable]:
"""返回 (out_lines, feed_line)。feed_line(line) 累积当前事件的若干行,
遇到空行表示事件结束,把整事件(清洗后)推入 out_lines。状态跨 chunk 存活。
"""
out: list[bytes] = []
evt: list[bytes] = []
def feed_line(ln: bytes):
if ln == b"":
if evt:
out.append(_process_event_lines(evt))
evt.clear()
else:
evt.append(ln)
def feed_all(lis: list[bytes]):
for ln in lis:
feed_line(ln)
return out, feed_all
# 全局唯一 sink(out_buf 累积所有已完整、待消费的事件)
out_buf, feed_all = _make_sink()
def _emit_sink() -> list[bytes] | None:
"""非阻塞:若 out_buf 非空则记录统计并清空,返回待 yield 的事件列表。"""
if not out_buf:
return None
batch = list(out_buf) # 快照后清空共享缓冲(不能直接引用再 clear,会清掉 batch)
out_buf.clear()
_record_stats(batch)
return batch
def _drain() -> list[bytes]:
"""把 out_buf 里所有事件经 reasoning 合并器后转成待转发列表。"""
out: list[bytes] = []
for evt in _emit_sink() or []:
if evt is not None:
for e in coal.feed(evt):
if e is not None:
forwarded_parts.append(e)
out.append(e)
return out
def _record_stats(events: list[bytes | None]):
nonlocal finish_reason, saw_filter
for cleaned in events:
if not cleaned:
continue
text_repr = cleaned.decode("utf-8", "replace")
# 统计用:从清洗后的 event 里解析(finish / tool_calls / usage)
for el in cleaned.split(b"\n"):
s = el.lstrip()
if not s.startswith(b"data:"):
continue
d = s[5:].lstrip()
if d == b"[DONE]":
continue
try: