Files
hack/tools/scripts/newapi_client.py
T

147 lines
5.7 KiB
Python

#!/usr/bin/env python3
"""Shared NewAPI client — used by all channel management scripts.
Centralizes HTTP requests, error handling, retries, and channel CRUD.
Secrets are read from environment variables with fallback to the
hardcoded defaults (which match the current setup).
"""
import json
import os
import time
import urllib.request
import urllib.error
from concurrent.futures import ThreadPoolExecutor, as_completed
# ── Config ───────────────────────────────────────────────────────
BASE = os.environ.get("NEWAPI_BASE", "https://ai.bor.nomsg.cn")
ADMIN_KEY = os.environ.get("NEWAPI_ADMIN_KEY", "CIoMSPSBNLt7rn8UE+qQDjjh9TNoeqg=")
HEADERS = {
"Authorization": f"Bearer {ADMIN_KEY}",
"Content-Type": "application/json",
}
# Channel type constants from NewAPI source
TYPE_OPENAI = 1
TYPE_CUSTOM = 8
TYPE_ANTHROPIC = 14
TYPE_ZHIPU_V4 = 26
TYPE_MINIMAX = 35
TYPE_OLLAMA = 37
TYPE_VOLC_ENGINE = 45
# ── Core HTTP ────────────────────────────────────────────────────
def _req(method, path, data=None, timeout=60, retries=2, backoff=1.0):
"""Send an HTTP request to the NewAPI server.
Returns the parsed JSON response dict. On HTTP errors the error
body is returned as a dict (preserving the original behaviour).
Network-level errors (DNS, timeout, connection refused) trigger a
retry with exponential backoff, then raise.
"""
url = f"{BASE}{path}"
body = json.dumps(data).encode() if data else None
last_exc = None
for attempt in range(retries + 1):
req = urllib.request.Request(url, data=body, headers=HEADERS, method=method)
try:
with urllib.request.urlopen(req, timeout=timeout) as resp:
raw = resp.read()
return json.loads(raw) if raw else {}
except urllib.error.HTTPError as e:
# Server responded with an HTTP error status — parse body
raw = e.read()
try:
return json.loads(raw) if raw else {}
except json.JSONDecodeError:
return {"success": False, "message": f"HTTP {e.code}: {raw[:200]}"}
except (urllib.error.URLError, TimeoutError, ConnectionError) as e:
last_exc = e
if attempt < retries:
time.sleep(backoff * (attempt + 1))
continue
return {"success": False, "message": f"Network error: {e}"}
except json.JSONDecodeError as e:
return {"success": False, "message": f"Bad JSON: {e}"}
return {"success": False, "message": f"Network error after {retries+1} attempts: {last_exc}"}
# ── Channel CRUD ─────────────────────────────────────────────────
def add_channel(channel, mode="single"):
"""Add a channel. Returns (success, message)."""
payload = {"mode": mode, "channel": channel}
data = _req("POST", "/api/channel/", payload)
return data.get("success", False), data.get("message", "")
def update_channel(cid, **fields):
"""Update a channel by ID. Returns (success, message)."""
payload = {"id": cid, **fields}
data = _req("PUT", "/api/channel/", payload)
return data.get("success", False), data.get("message", "")
def delete_channel(cid):
"""Delete a channel by ID. Returns (success, message)."""
data = _req("DELETE", f"/api/channel/{cid}")
return data.get("success", False), data.get("message", "")
def get_channel(cid):
"""Get a single channel by ID. Returns the channel dict."""
data = _req("GET", f"/api/channel/{cid}")
return data.get("data", {})
def list_channels(page_size=200):
"""List all channels. Returns a list of channel dicts (paginates)."""
# The server silently caps page_size at 100 regardless of what we
# request, so paginate using the actual returned count (not the
# requested one) to decide when we've reached the last page.
page = 100
all_items = []
p = 0
while True:
data = _req("GET", f"/api/channel/?p={p}&page_size={page}")
items = data.get("data", {}).get("items", [])
if not items:
break
all_items.extend(items)
if len(items) < page:
break
p += 1
return all_items
def test_channel(cid, timeout=60):
"""Test a channel by ID. Returns (success, message, time)."""
data = _req("GET", f"/api/channel/test/{cid}", timeout=timeout)
return data.get("success", False), data.get("message", ""), data.get("time", 0)
# ── Batch helpers ────────────────────────────────────────────────
def test_channels_concurrent(channel_ids, max_workers=8):
"""Test multiple channels concurrently.
Returns a list of (cid, success, message, elapsed) tuples.
"""
results = []
with ThreadPoolExecutor(max_workers=max_workers) as pool:
futures = {pool.submit(test_channel, cid): cid for cid in channel_ids}
for future in as_completed(futures):
cid = futures[future]
try:
success, msg, t = future.result()
except Exception as e:
success, msg, t = False, str(e), 0
results.append((cid, success, msg, t))
# Sort by channel ID for stable output
results.sort(key=lambda r: r[0])
return results
def list_channel_names():
"""Return a set of all existing channel names (for dedup checks)."""
return {ch["name"] for ch in list_channels()}