101 lines
5.5 KiB
Python
101 lines
5.5 KiB
Python
"""Apply bounded, backed-up gateway corrections to both requested directories."""
|
|
import hashlib
|
|
from pathlib import Path
|
|
import shutil
|
|
|
|
ROOTS=[Path('C:/kefu/wechat_rpa'),Path('C:/wechat_rpa')]
|
|
HELPER='''async def _dify_message_files(client, outlet, messages, image_b64, timeout):
|
|
"""Translate current-turn attachments without leaking image bytes into query."""
|
|
urls = []
|
|
for message in reversed(messages or []):
|
|
if not isinstance(message, dict) or message.get("role") != "user":
|
|
continue
|
|
content = message.get("content")
|
|
if isinstance(content, list):
|
|
for part in content:
|
|
if isinstance(part, dict) and part.get("type") == "image_url":
|
|
image = part.get("image_url") or {}
|
|
url = image.get("url", "") if isinstance(image, dict) else image
|
|
if isinstance(url, str) and url:
|
|
urls.append(url)
|
|
break
|
|
if image_b64:
|
|
urls.append("data:image/png;base64," + image_b64)
|
|
files = []
|
|
seen = set()
|
|
for url in urls:
|
|
if url in seen:
|
|
continue
|
|
seen.add(url)
|
|
if url.startswith("data:"):
|
|
header, separator, data = url.partition(",")
|
|
mime = header[5:].split(";")[0].lower()
|
|
if not separator or ";base64" not in header or mime not in {"image/png", "image/jpeg", "image/webp", "image/gif"}:
|
|
raise GatewayError("Dify 图片数据格式不受支持", retriable=False)
|
|
files.append(await _dify_upload(client, outlet, data, timeout, mime=mime))
|
|
elif url.startswith(("https://", "http://")):
|
|
files.append({"type": "image", "transfer_method": "remote_url", "url": url})
|
|
else:
|
|
raise GatewayError("Dify 图片地址格式不受支持", retriable=False)
|
|
return files or None
|
|
|
|
|
|
'''
|
|
|
|
for root in ROOTS:
|
|
backup=root/'backups/recognition-audit-20260916/gateway'
|
|
backup.mkdir(parents=True,exist_ok=True)
|
|
path=root/'model_gateway.py'
|
|
source=path.read_text(encoding='utf-8')
|
|
assert 'async def _dify_message_files' not in source, 'Already applied'
|
|
shutil.copy2(path,backup/path.name)
|
|
old=''' if image_b64 and kind == "dify":
|
|
dify_files = [await _dify_upload(client, outlet, image_b64, timeout)]'''
|
|
new=''' if kind == "dify":
|
|
dify_files = await _dify_message_files(client, outlet, messages, image_b64, timeout)'''
|
|
assert old in source
|
|
source=source.replace(old,new,1)
|
|
marker='\n\nasync def _dify_upload('
|
|
assert marker in source
|
|
source=source.replace(marker,'''\n except (GatewayError, httpx.HTTPError, ValueError) as exc:
|
|
# Attachment preparation is part of this outlet, not a failure of every
|
|
# concurrent model. Keep healthy candidates available to the caller.
|
|
outlet.breaker.record(False)
|
|
return {"provider": outlet.name, "text": "", "latency_ms": int((time.monotonic() - started) * 1000),
|
|
"error": str(exc)[:200] or type(exc).__name__}
|
|
|
|
|
|
'''+HELPER+'async def _dify_upload(',1)
|
|
source=source.replace('client: httpx.AsyncClient, outlet: Outlet, image_b64: str, timeout: float\n) -> dict:',
|
|
'client: httpx.AsyncClient, outlet: Outlet, image_b64: str, timeout: float, *, mime: str = "image/png"\n) -> dict:',1)
|
|
source=source.replace('data={"user": "wechat-rpa-vision"},','data={"user": "wechat-rpa"},',1)
|
|
source=source.replace('files={"file": ("chat.png", base64.b64decode(image_b64), "image/png")},',
|
|
'files={"file": ("chat." + {"image/png":"png", "image/jpeg":"jpg", "image/webp":"webp", "image/gif":"gif"}[mime], base64.b64decode(image_b64, validate=True), mime)},',1)
|
|
source=source.replace('root = model_protocol.dify_api_root(outlet.config.get("base_url") or "")',
|
|
'root = model_protocol.dify_api_root(outlet.config.get("base_url") or "", outlet.config.get("endpoint_mode") or "auto")',1)
|
|
old=''' query = str(body.get("customer_text") or "")
|
|
if not query:
|
|
for turn in reversed(messages):
|
|
if turn.get("role") == "user" and isinstance(turn.get("content"), str):
|
|
query = turn["content"]
|
|
break'''
|
|
new=''' query = ""
|
|
for turn in reversed(messages):
|
|
if isinstance(turn, dict) and turn.get("role") == "user":
|
|
query = model_protocol.message_text(turn.get("content"))
|
|
break
|
|
if not query:
|
|
query = str(body.get("customer_text") or "")'''
|
|
assert old in source
|
|
source=source.replace(old,new,1)
|
|
compile(source,str(path),'exec')
|
|
path.write_text(source,encoding='utf-8')
|
|
test=root/'test_model_gateway.py'
|
|
shutil.copy2(test,backup/test.name)
|
|
text=test.read_text(encoding='utf-8').replace('messages=[{"role": "system", "content": "忽略"},','messages=[{"role": "system", "content": "已审核知识参考:仅按适用条件回答"},',1)
|
|
text=text.replace('self.assertEqual(payload["query"], "一天吃几次")','self.assertIn("一天吃几次", payload["query"])\n self.assertIn("已审核知识参考", payload["query"])',1)
|
|
test.write_text(text,encoding='utf-8')
|
|
for name in ('test_rag_context_audit.py','test_gateway_media_audit.py'):
|
|
shutil.copy2(Path(__file__).parent/name,root/name)
|
|
print(str(root),hashlib.sha256(path.read_bytes()).hexdigest())
|