166 lines
9.4 KiB
Python
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')
|