Files
kefu/deploy/token-usage-20260917/implement_collection.py
T
2026-09-21 10:34:06 +08:00

166 lines
9.4 KiB
Python

from pathlib import Path
import shutil
root = Path('C:/kefu/wechat_rpa')
backup = Path(__file__).parent / 'before-collection'
backup.mkdir(exist_ok=True)
for name in ('model_protocol.py', 'model_gateway.py', 'knowledge_worker.py'):
if not (backup / name).exists():
shutil.copy2(root / name, backup / name)
def edit(name, old, new, count=1):
path = root / name
text = path.read_text(encoding='utf-8')
if text.count(old) != count:
raise RuntimeError(f'{name}: expected {count} anchors, got {text.count(old)}: {old[:100]}')
path.write_text(text.replace(old, new), encoding='utf-8', newline='\n')
parser = '''def parse_usage(kind: str, data: object, base_url: str = "") -> dict:
"""Normalize reported tokens; missing/invalid counters stay unknown (None).
OpenAI cache/reasoning details are included in their parent counters. Claude
input_tokens excludes cache reads/writes, so include both once. Dify reports
blocking chat usage under metadata. Never estimate tokens from characters.
"""
def obj(value):
return value if isinstance(value, dict) else {}
def counter(value):
# bool is an int; reject it and fractional/negative/oversized values.
if isinstance(value, bool):
return None
if isinstance(value, str) and value.isascii() and value.isdecimal() and len(value) <= 15:
value = int(value)
return value if isinstance(value, int) and 0 <= value <= 10**12 else None
kind = detect_kind(kind, base_url)
response = obj(data)
usage = obj(obj(response.get("metadata")).get("usage")) if kind == "dify" else obj(response.get("usage"))
if kind == "claude":
input_tokens = counter(usage.get("input_tokens"))
cache_read = counter(usage.get("cache_read_input_tokens", 0))
cache_write = counter(usage.get("cache_creation_input_tokens", 0))
if all(value is not None for value in (input_tokens, cache_read, cache_write)):
input_tokens += cache_read + cache_write
else:
input_tokens = None
output_tokens = counter(usage.get("output_tokens"))
cached = counter(usage.get("cache_read_input_tokens"))
reasoning = counter(obj(usage.get("output_tokens_details")).get("thinking_tokens"))
else:
input_tokens = counter(usage.get("prompt_tokens", usage.get("input_tokens")))
output_tokens = counter(usage.get("completion_tokens", usage.get("output_tokens")))
cached = counter(obj(usage.get("prompt_tokens_details", usage.get("input_tokens_details"))).get("cached_tokens"))
reasoning = counter(obj(usage.get("completion_tokens_details", usage.get("output_tokens_details"))).get("reasoning_tokens"))
total = counter(usage.get("total_tokens"))
if total is None and input_tokens is not None and output_tokens is not None:
total = input_tokens + output_tokens
return {"input_tokens": input_tokens, "output_tokens": output_tokens,
"total_tokens": total, "cached_input_tokens": cached, "reasoning_tokens": reasoning}
'''
edit('model_protocol.py', 'def image_message(', parser + 'def image_message(')
edit('model_gateway.py', 'import admin_backend\n', 'import admin_backend\nfrom model_usage import UsageRecorder\n')
edit('model_gateway.py', 'def __init__(self, message: str, *, retriable: bool = False):\n super().__init__(message)\n self.retriable = retriable',
'def __init__(self, message: str, *, retriable: bool = False, response_data=None):\n super().__init__(message)\n self.retriable = retriable\n self.response_data = response_data')
edit('model_gateway.py', ''' if response.status_code in RETRIABLE_STATUS:
raise GatewayError(
f"上游返回 {response.status_code}", retriable=True
)
if response.status_code >= 400:
raise GatewayError(
f"上游返回 {response.status_code}: {response.text[:200]}", retriable=False
)
try:
return response.json()
except ValueError as exc:
raise GatewayError("上游响应不是合法 JSON", retriable=False) from exc
''', ''' try:
data = response.json()
except ValueError:
data = None
if response.status_code in RETRIABLE_STATUS:
raise GatewayError(
f"上游返回 {response.status_code}", retriable=True, response_data=data
)
if response.status_code >= 400:
raise GatewayError(
f"上游返回 {response.status_code}: {response.text[:200]}", retriable=False, response_data=data
)
if data is None:
raise GatewayError("上游响应不是合法 JSON", retriable=False)
return data
''')
edit('model_gateway.py', ' deadline: float = SINGLE_CALL_TIMEOUT,\n',
' deadline: float = SINGLE_CALL_TIMEOUT,\n usage_recorder: UsageRecorder | None = None,\n usage_role: str = "answer",\n')
edit('model_gateway.py', ''' for attempt in range(1, MAX_ATTEMPTS + 1):
try:''', ''' for attempt in range(1, MAX_ATTEMPTS + 1):
data = None
status = "error"
attempt_started = time.monotonic()
try:''')
edit('model_gateway.py', ''' outlet.breaker.record(True)
return {''', ''' status = "success"
outlet.breaker.record(True)
return {''')
edit('model_gateway.py', ''' except GatewayError as exc:
last = exc''', ''' except GatewayError as exc:
data = exc.response_data
last = exc''')
edit('model_gateway.py', ''' await asyncio.sleep(0.35 * attempt + secrets.randbelow(200) / 1000.0)
except ValueError''', ''' except ValueError''')
edit('model_gateway.py', ''' last = exc
break
outlet.breaker.record(False)''', ''' last = exc
break
finally:
if usage_recorder is not None:
usage_recorder.capture(
outlet, data, attempt=attempt, role=usage_role, status=status,
latency_ms=int((time.monotonic() - attempt_started) * 1000),
)
await asyncio.sleep(0.35 * attempt + secrets.randbelow(200) / 1000.0)
outlet.breaker.record(False)''')
edit('model_gateway.py', ' second: str = "",\n) -> dict:',
' second: str = "",\n *, usage_recorder: UsageRecorder | None = None,\n) -> dict:')
edit('model_gateway.py', ' deadline=JUDGE_DEADLINE,\n',
' deadline=JUDGE_DEADLINE,\n usage_recorder=usage_recorder, usage_role="judge",\n')
# Every return (including 502) flushes the same recorder outside upstream deadlines.
path = root / 'model_gateway.py'
text = path.read_text(encoding='utf-8')
start = text.index(' started = time.monotonic()\n knowledge_result = None')
end = text.index('\n return app', start)
block = text[start:end]
block = block.replace('deadline=ANSWER_DEADLINE,', 'deadline=ANSWER_DEADLINE,\n usage_recorder=usage_recorder,')
block = block.replace('usable[1]["text"] if len(usable) > 1 else "",',
'usable[1]["text"] if len(usable) > 1 else "",\n usage_recorder=usage_recorder,')
block = block.replace('str(body.get("task_id") or "")', 'usage_recorder.context["task_id"]')
text = text[:start] + ''' async with UsageRecorder(
database, tenant_id=str(account["tenant_id"]), desktop_account_id=int(account["id"]),
task_id=str(body.get("task_id") or ""), purpose=purpose,
) as usage_recorder:
''' + ''.join(' ' + line if line.strip() else line for line in block.splitlines(keepends=True)) + text[end:]
path.write_text(text, encoding='utf-8', newline='\n')
edit('knowledge_worker.py', 'def enhance(candidate, database, options=None):',
'def enhance(candidate, database, options=None, *, tenant_id="default", task_id=""):')
edit('knowledge_worker.py', ' from model_gateway import call_outlet\n',
' from model_gateway import call_outlet\n from model_usage import UsageRecorder\n')
edit('knowledge_worker.py', ''' async def run():
async with httpx.AsyncClient() as client:
return await call_outlet(client, outlet, [{"role": "user", "content": prompt}],
deadline=deadline, max_tokens=4096, temperature=0.1)
''', ''' async def run():
async with UsageRecorder(database, tenant_id=tenant_id, task_id=task_id,
purpose="knowledge") as recorder:
async with httpx.AsyncClient() as client:
return await call_outlet(client, outlet, [{"role": "user", "content": prompt}],
deadline=deadline, max_tokens=4096, temperature=0.1,
usage_recorder=recorder)
''')
edit('knowledge_worker.py', 'candidate, input_chars, output_chars = enhance(candidate, self.store.database, options)',
'candidate, input_chars, output_chars = enhance(\n candidate, self.store.database, options, tenant_id=job["tenant_id"], task_id=job["id"])')
print('Updated collection in model_protocol.py, model_gateway.py, knowledge_worker.py')