同步客户绑定与登录恢复优化

This commit is contained in:
chen qiang
2026-10-08 19:51:38 +08:00
parent 2a81f5668e
commit 3e25ca4696
12 changed files with 896 additions and 152 deletions
+317 -140
View File
@@ -28,6 +28,17 @@ import secrets
from typing import Optional, Dict, Any, List, Tuple
import argparse
import copy
from contextlib import redirect_stdout
try:
from .runtime_state import ConfigSnapshot, file_lock, read_config, update_config
except ImportError:
from runtime_state import ConfigSnapshot, file_lock, read_config, update_config
EXECUTION_CONTEXT = None
JSON_OUTPUT = False
JSON_STREAM = None
TARGET_KEYS = ("client_code", "client_title", "server_id", "transit_url")
from datetime import datetime, timezone
# 修复 Windows 控制台 UTF-8 输出
@@ -40,7 +51,7 @@ if sys.platform == 'win32':
# ============================================================
TRANSIT_URL = "https://ts.henlo.net"
AUTH_TOKEN = ""
DEFAULT_SERVER_ID = "server-001"
DEFAULT_SERVER_ID = ""
DEFAULT_TIMEOUT = 60 # 默认超时(秒)
MAX_RETRIES = 3 # 最大重试次数
RETRY_DELAY = 2 # 重试延迟(秒)
@@ -340,34 +351,62 @@ def get_script_dir():
def get_config():
"""Load local runtime config, creating it from the checked-in template when absent."""
"""Read fresh credentials while preserving the command's fixed target."""
cfg_path = os.path.join(get_script_dir(), "config.json")
template_path = os.path.join(get_script_dir(), "config.template.json")
for path in (cfg_path, "config.json"):
try:
with open(path, "r", encoding="utf-8-sig") as f:
return json.load(f)
cfg = read_config(path)
if path != cfg_path:
cfg = update_config(cfg_path, dict(cfg), initialize=True)
break
except FileNotFoundError:
continue
except json.JSONDecodeError as exc:
logger.warning("Invalid local config %s: %s", path, exc)
break
cfg = {"transit_url": TRANSIT_URL, "server_id": "", "access_token": "", "expires_at": "", "user_name": "", "client_code": "", "client_title": "", "client_list": []}
try:
with open(template_path, "r", encoding="utf-8-sig") as f:
cfg.update(json.load(f))
except (FileNotFoundError, json.JSONDecodeError):
pass
save_config(cfg)
else:
template = os.path.join(get_script_dir(), "config.template.json")
try:
template_cfg = read_config(template)
except FileNotFoundError:
template_cfg = {"transit_url": TRANSIT_URL, "server_id": "", "access_token": "", "expires_at": "", "user_name": "", "client_code": "", "client_title": "", "client_list": []}
cfg = update_config(cfg_path, dict(template_cfg), initialize=True)
if EXECUTION_CONTEXT is not None:
cfg.update(EXECUTION_CONTEXT)
# Context fields are a view, not persistent edits.
cfg.original = copy.deepcopy(dict(cfg))
return cfg
def save_config(cfg: dict):
"""保存配置到脚本同目录的 config.json"""
cfg_path = os.path.join(get_script_dir(), "config.json")
with open(cfg_path, "w", encoding="utf-8") as f:
json.dump(cfg, f, ensure_ascii=False, indent=2)
logger.info(f"配置已保存到 {cfg_path}")
ignored = TARGET_KEYS if EXECUTION_CONTEXT is not None else ()
update_config(os.path.join(get_script_dir(), "config.json"), cfg, ignored)
def _authenticated_request(method, url, **kwargs):
"""Recover one rejected session; never replay transport errors here."""
request = getattr(requests, method)
response = request(url, **kwargs)
headers = kwargs.get("headers", {})
authorization = headers.get("Authorization", "")
token = authorization[7:] if authorization.startswith("Bearer ") else ""
if (response.status_code != 401 or not token or url.endswith("/logout") or
"/api/admin/" in url or url.endswith("/api/device/challenge") or url.endswith("/api/device/login")):
return response
transit = get_transit_url().rstrip("/")
if not url.startswith(transit + "/api/"):
return response
cfg = get_config()
if not _auto_device_login(cfg, rejected_token=token):
return response
response.close()
kwargs["headers"] = dict(headers, Authorization="Bearer " + cfg["access_token"])
return request(url, **kwargs)
def _http_get(url, **kwargs):
return _authenticated_request("get", url, **kwargs)
def _http_post(url, **kwargs):
return _authenticated_request("post", url, **kwargs)
def get_client_code(client: Dict[str, Any]) -> str:
@@ -444,27 +483,11 @@ def refresh_client_cache(cfg: Optional[Dict[str, Any]] = None, quiet: bool = Fal
clients: List[Dict[str, Any]] = normalize_client_list(cfg.get("client_list", []))
try:
resp = requests.get(f"{transit_url}/api/me", headers=make_headers(), timeout=10)
resp = _http_get(f"{transit_url}/api/clients", headers=make_headers(), timeout=10)
if resp.status_code == 401:
if not quiet:
print("❌ 中转机登录态已失效,请重新登录")
return clients
resp.raise_for_status()
me = resp.json()
if me.get("expires_at"):
cfg["expires_at"] = me.get("expires_at")
if me.get("user", {}).get("name"):
cfg["user_name"] = me.get("user", {}).get("name")
except requests.exceptions.RequestException as e:
if not quiet:
print(f"⚠️ 从 /api/me 刷新 client 失败: {e}")
try:
resp = requests.get(f"{transit_url}/api/clients", headers=make_headers(), timeout=10)
if resp.status_code == 401:
if not quiet:
print("❌ 中转机登录态已失效,请重新登录")
return clients
print("❌ 中转机登录态已失效,请重新登录", file=sys.stderr)
return []
resp.raise_for_status()
data = resp.json()
# /api/clients is authoritative for the current authorization snapshot.
@@ -472,10 +495,17 @@ def refresh_client_cache(cfg: Optional[Dict[str, Any]] = None, quiet: bool = Fal
clients = normalize_client_list(data.get("clients", []))
except requests.exceptions.RequestException as e:
if not quiet:
print(f"⚠️ 从 /api/clients 刷新 client 失败: {e}")
return clients
print(f"⚠️ 从 /api/clients 刷新 client 失败: {e}", file=sys.stderr)
return []
if sync_client_cache(cfg, clients):
# A 401 recovery may have renewed credentials during the request.
fresh = get_config()
for key in ("access_token", "expires_at", "user_name"):
if key in fresh:
cfg[key] = fresh[key]
if isinstance(cfg, ConfigSnapshot):
cfg.original[key] = copy.deepcopy(fresh[key])
save_config(cfg)
return clients
@@ -538,7 +568,7 @@ def ts_api_call(method: str, params: dict, timeout: int = 15) -> Dict[str, Any]:
payload = {"method": method, "params": params}
try:
logger.info(f"调用 BOS 接口: {method}")
resp = requests.post(BOS_API_URL, json=payload, timeout=timeout)
resp = _http_post(BOS_API_URL, json=payload, timeout=timeout)
resp.raise_for_status()
result = resp.json()
logger.info(f"BOS 接口响应: code={result.get('code')}, success={result.get('success')}")
@@ -585,7 +615,7 @@ def get_server_status(client_code: str = "") -> Dict[str, Any]:
if not client_code:
return {"success": False, "error": "clientCode is required"}
try:
resp = requests.get(
resp = _http_get(
f"{transit_url}/api/server_status",
params={"client_code": client_code},
headers=make_headers(),
@@ -719,7 +749,7 @@ def inspection_report(client_code: str = "", timeout: int = 90, report_type: str
"timeout": timeout,
}
try:
resp = requests.post(
resp = _http_post(
f"{transit_url}/api/inspection_report",
json=payload,
headers=make_headers(),
@@ -754,7 +784,7 @@ def inspection_report_latest(client_code: str = "", refresh: bool = False) -> Di
if refresh:
params["refresh"] = "true"
try:
resp = requests.get(
resp = _http_get(
f"{transit_url}/api/inspection_report/latest",
params=params,
headers=make_headers(),
@@ -781,7 +811,7 @@ def inspection_report_get(report_id: str) -> Dict[str, Any]:
if not report_id:
return {"success": False, "error": "report id is required"}
try:
resp = requests.get(
resp = _http_get(
f"{transit_url}/api/inspection_report/{report_id}",
headers=make_headers(),
timeout=DEFAULT_TIMEOUT + 5,
@@ -811,7 +841,7 @@ def awr_status(client_code: str = "") -> Dict[str, Any]:
if not client_code:
return {"success": False, "error": "clientCode is required"}
try:
resp = requests.get(
resp = _http_get(
f"{str(cfg.get('transit_url') or TRANSIT_URL).rstrip('/')}/api/awr/status",
params={"client_code": client_code},
headers=make_headers(),
@@ -849,7 +879,7 @@ def awr_list(client_code: str = "") -> Dict[str, Any]:
if not client_code:
return {"success": False, "error": "clientCode is required"}
try:
resp = requests.get(
resp = _http_get(
f"{str(cfg.get('transit_url') or TRANSIT_URL).rstrip('/')}/api/awr/list",
params={"client_code": client_code},
headers=make_headers(),
@@ -897,7 +927,7 @@ def download_awr_report(client_code: str, date: str, output_path: str = "") -> D
return {"success": False, "error": str(e)}
total = 0
try:
with requests.get(
with _http_get(
f"{str(cfg.get('transit_url') or TRANSIT_URL).rstrip('/')}/api/awr/download",
params={"client_code": client_code, "date": date},
headers=make_headers(),
@@ -994,7 +1024,7 @@ def _log_api(endpoint: str, payload: Dict[str, Any], timeout: int = 120) -> Dict
body["client_code"] = client_code
body["timeout"] = max(1, min(int(timeout or 120), 300))
try:
resp = requests.post(
resp = _http_post(
f"{str(cfg.get('transit_url') or TRANSIT_URL).rstrip('/')}/api/log/{endpoint}",
json=body,
headers=make_headers(),
@@ -1683,7 +1713,7 @@ def download_inspection_report_html(result: Dict[str, Any], path: str = "") -> s
last_error = ""
for attempt in range(5):
try:
response = requests.get(html_url, headers=make_headers(), timeout=30)
response = _http_get(html_url, headers=make_headers(), timeout=30)
if response.status_code != 404:
response.raise_for_status()
break
@@ -1813,7 +1843,7 @@ def agent_update(client_code: str = "", timeout: int = 300) -> Dict[str, Any]:
headers = make_headers()
headers["X-Upgrade-Admin-Token"] = upgrade_token
try:
resp = requests.post(
resp = _http_post(
f"{transit_url}/api/admin/agent_update/trigger",
json=payload,
headers=headers,
@@ -1872,7 +1902,7 @@ def oracle_ops(item: str, client_code: str = "", minutes: int = 60, top_n: int =
"timeout": timeout,
}
try:
resp = requests.post(f"{transit_url}/api/oracle_ops", json=payload, headers=make_headers(), timeout=timeout + 5)
resp = _http_post(f"{transit_url}/api/oracle_ops", json=payload, headers=make_headers(), timeout=timeout + 5)
if resp.status_code in (400, 401, 403, 503):
try:
return {"success": False, "error": resp.json().get("error", resp.text)}
@@ -1900,7 +1930,7 @@ def oracle_ops_report(client_code: str = "", items=None, minutes: int = 60, top_
"per_item_timeout": per_item_timeout,
}
try:
resp = requests.post(
resp = _http_post(
f"{transit_url}/api/oracle_ops_report",
json=payload,
headers=make_headers(),
@@ -2018,6 +2048,11 @@ def cmd_sys_functions(schema: str = "", as_json: bool = False):
def cmd_login(secret_key: str, client_code: str = ""):
with file_lock(os.path.join(get_script_dir(), "config.json.auth.lock")):
return _cmd_login_unlocked(secret_key, client_code)
def _cmd_login_unlocked(secret_key: str, client_code: str = ""):
"""
登录中转机并选择 client
@@ -2030,7 +2065,7 @@ def cmd_login(secret_key: str, client_code: str = ""):
transit_url = cfg.get("transit_url", TRANSIT_URL)
try:
resp = requests.post(
resp = _http_post(
f"{transit_url}/api/login",
json={"secret_key": secret_key, "client_code": client_code},
timeout=20,
@@ -2129,7 +2164,7 @@ def cmd_device_register(device_name: str = ""):
mac_hash, mac_masked = _device_identity()
body = {"device_id": key["device_id"], "device_name": device_name or platform.node() or "Trusted device", "computer_name": platform.node(), "os_name": platform.platform(), "key_algorithm": "ED25519", "public_key": key["public_key"], "key_fingerprint": key["key_fingerprint"], "mac_hash": mac_hash, "mac_masked": mac_masked, "description": "Registered by oracle-jump-query skill"}
try:
resp = requests.post(f"{cfg.get('transit_url', TRANSIT_URL)}/api/device/register", json=body, headers=make_headers(), timeout=20)
resp = _http_post(f"{cfg.get('transit_url', TRANSIT_URL)}/api/device/register", json=body, headers=make_headers(), timeout=20)
data = resp.json()
except (requests.RequestException, ValueError) as exc:
print(f"❌ 可信设备注册失败: {exc}")
@@ -2144,6 +2179,18 @@ def cmd_device_register(device_name: str = ""):
def cmd_device_login(client_code: str = ""):
started_token = get_config().get("access_token", "")
with file_lock(os.path.join(get_script_dir(), "config.json.auth.lock")):
cfg = get_config()
if cfg.get("access_token") and cfg["access_token"] != started_token and not is_token_expired(cfg):
print("✅ 已复用本机有效 token")
if client_code:
cmd_switch(client_code)
return
return _cmd_device_login_unlocked(client_code)
def _cmd_device_login_unlocked(client_code: str = ""):
"""Use the locally stored Ed25519 key to perform trusted-device login."""
try:
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
@@ -2155,14 +2202,14 @@ def cmd_device_login(client_code: str = ""):
transit_url = (get_config() or {}).get("transit_url", TRANSIT_URL)
identity = {"device_id": key["device_id"], "key_fingerprint": key["key_fingerprint"]}
try:
challenge_resp = requests.post(f"{transit_url}/api/device/challenge", json=identity, timeout=20)
challenge_resp = _http_post(f"{transit_url}/api/device/challenge", json=identity, timeout=20)
challenge_data = challenge_resp.json()
if not challenge_data.get("success"):
print(f"❌ 获取可信设备挑战失败: {challenge_data.get('error', challenge_resp.text)}")
return
challenge = challenge_data["challenge"]
signature = base64.b64encode(private.sign(challenge.encode("utf-8"))).decode("ascii")
resp = requests.post(f"{transit_url}/api/device/login", json={**identity, "challenge": challenge, "signature": signature, "client_code": client_code}, timeout=20)
resp = _http_post(f"{transit_url}/api/device/login", json={**identity, "challenge": challenge, "signature": signature, "client_code": client_code}, timeout=20)
data = resp.json()
except (requests.RequestException, ValueError) as exc:
print(f"❌ 可信设备登录失败: {exc}")
@@ -2181,83 +2228,47 @@ def cmd_device_login(client_code: str = ""):
def cmd_logout():
with file_lock(os.path.join(get_script_dir(), "config.json.auth.lock")):
return _cmd_logout_unlocked()
def _cmd_logout_unlocked():
"""登出,清除登录信息"""
cfg = get_config() or {}
token = cfg.get("access_token", "")
transit_url = cfg.get("transit_url", TRANSIT_URL)
if token:
try:
requests.post(f"{transit_url}/api/logout", headers=make_headers(), timeout=10)
_http_post(f"{transit_url}/api/logout", headers=make_headers(), timeout=10)
except Exception:
pass
# 保留 transit_url 和 server_id,清除登录相关
keep_keys = {"transit_url", "server_id", "default_schema"}
new_cfg = {k: v for k, v in cfg.items() if k in keep_keys}
save_config(new_cfg)
for key in list(cfg):
if key not in keep_keys:
del cfg[key]
save_config(cfg)
print("✅ 已登出")
def cmd_status():
"""显示当前登录状态和选中的 client"""
config = get_config() or {}
transit_url = config.get("transit_url", TRANSIT_URL)
"""Validate the session and Agent state with a single clients request."""
cfg = get_config()
if not ensure_logged_in(cfg):
return
clients = refresh_client_cache(cfg)
cfg = get_config()
print("当前状态:")
print("-" * 40)
if not config.get("access_token"):
print(" 未登录,请使用 login <secretKey> [clientCode] 登录")
print("-" * 40)
return
print(" 用户:" + cfg.get("user_name", "未知"))
print(" client:" + cfg.get("client_code", cfg.get("server_id", "")))
print(" 过期时间:" + cfg.get("expires_at", ""))
for client in clients:
code = get_client_code(client)
marker = " <--" if code == cfg.get("client_code") else ""
online = "在线" if client.get("online") else "离线"
auth = "已授权" if client.get("authorized") else "未授权"
print(" %s (%s) %s | %s%s" % (code, get_client_title(client), online, auth, marker))
if is_token_expired(config):
print(" 登录已过期,请重新执行 login <secretKey> [clientCode]")
print(f" 过期时间:{config.get('expires_at', '')}")
print("-" * 40)
return
try:
resp = requests.get(f"{transit_url}/api/me", headers=make_headers(), timeout=5)
if resp.status_code == 401:
print(" 中转机登录态已失效,请重新登录")
print("-" * 40)
return
resp.raise_for_status()
me = resp.json()
user = me.get("user", {})
client_list = me.get("client_list", [])
if sync_client_cache(config, client_list):
if me.get("expires_at"):
config["expires_at"] = me.get("expires_at")
if user.get("name"):
config["user_name"] = user.get("name")
save_config(config)
print(f" 用户:{user.get('name', config.get('user_name', '未知'))}")
print(f" client:{config.get('client_code', '未选择')} ({config.get('client_title', '')})")
print(f" 过期时间:{me.get('expires_at', config.get('expires_at', ''))}")
if client_list:
print(f" 可用 client:")
for cl in client_list:
marker = " <--" if cl["code"] == config.get("client_code") else ""
print(f" - {cl['code']} ({cl.get('title', '')}){marker}")
except Exception as e:
print(f" 无法从中转机获取状态: {e}")
# 显示 Agent 在线状态
try:
resp = requests.get(f"{transit_url}/api/clients", headers=make_headers(), timeout=5)
if resp.status_code == 200:
data = resp.json()
if sync_client_cache(config, data.get("clients", [])):
save_config(config)
current_code = config.get("client_code", "")
for cl in data.get("clients", []):
if cl.get("code") == current_code:
online = "在线" if cl.get("online") else "离线"
auth = "已授权" if cl.get("authorized") else "未授权"
print(f" Agent 状态:{online} | {auth}")
break
except Exception:
pass # 静默失败
print("-" * 40)
def cmd_switch(client_selector: str):
"""切换当前 client,支持 code 或 title/name 匹配。"""
@@ -2328,7 +2339,28 @@ def is_token_expired(cfg: dict) -> bool:
return now >= expires_at
def _auto_device_login(cfg: Dict[str, Any]) -> bool:
def _auto_device_login(cfg: Dict[str, Any], rejected_token=None) -> bool:
try:
with file_lock(os.path.join(get_script_dir(), "config.json.auth.lock")):
fresh = get_config()
token = fresh.get("access_token", "")
if token and not is_token_expired(fresh) and token != rejected_token:
cfg.update(fresh)
if isinstance(cfg, ConfigSnapshot):
cfg.original = copy.deepcopy(dict(cfg))
return True
if not _auto_device_login_unlocked(fresh):
return False
cfg.update(fresh)
if isinstance(cfg, ConfigSnapshot):
cfg.original = copy.deepcopy(dict(cfg))
return True
except (OSError, TimeoutError) as exc:
print("❌ 可信设备续登失败: " + str(exc), file=sys.stderr)
return False
def _auto_device_login_unlocked(cfg: Dict[str, Any]) -> bool:
"""Refresh the transit session with the locally approved trusted device."""
try:
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
@@ -2336,21 +2368,24 @@ def _auto_device_login(cfg: Dict[str, Any]) -> bool:
private = Ed25519PrivateKey.from_private_bytes(base64.b64decode(key["private_key"]))
identity = {"device_id": key["device_id"], "key_fingerprint": key["key_fingerprint"]}
transit_url = cfg.get("transit_url", TRANSIT_URL)
challenge_data = requests.post(f"{transit_url}/api/device/challenge", json=identity, timeout=20).json()
challenge_data = _http_post(f"{transit_url}/api/device/challenge", json=identity, timeout=20).json()
if not challenge_data.get("success"):
print("❌ 可信设备挑战失败: " + str(challenge_data.get("error", "unknown")), file=sys.stderr)
return False
challenge = challenge_data["challenge"]
signature = base64.b64encode(private.sign(challenge.encode("utf-8"))).decode("ascii")
data = requests.post(f"{transit_url}/api/device/login", json={**identity, "challenge": challenge, "signature": signature, "client_code": cfg.get("client_code", "")}, timeout=20).json()
data = _http_post(f"{transit_url}/api/device/login", json={**identity, "challenge": challenge, "signature": signature, "client_code": cfg.get("client_code", "")}, timeout=20).json()
if not data.get("success") or not data.get("access_token"):
print("❌ 可信设备续登失败: " + str(data.get("error", "missing token")), file=sys.stderr)
return False
current = data.get("current_client", {})
cfg.update({"access_token": data["access_token"], "expires_at": data.get("expires_at", ""), "user_name": data.get("user", {}).get("name", cfg.get("user_name", "")), "client_code": current.get("code", cfg.get("client_code", "")), "client_title": current.get("title", cfg.get("client_title", "")), "client_list": data.get("client_list", cfg.get("client_list", []))})
cfg["server_id"] = cfg.get("client_code", cfg.get("server_id", ""))
save_config(cfg)
print("✅ 中转 token 已通过可信设备自动续取")
print("✅ 中转 token 已通过可信设备自动续取", file=sys.stderr)
return True
except (ImportError, OSError, KeyError, ValueError, TypeError, requests.RequestException):
except (ImportError, OSError, KeyError, ValueError, TypeError, requests.RequestException) as exc:
print("❌ 可信设备续登失败: " + str(exc), file=sys.stderr)
return False
@@ -2360,11 +2395,11 @@ def ensure_logged_in(cfg: Optional[dict] = None) -> bool:
if _auto_device_login(cfg):
return True
if not cfg.get("access_token"):
print("❌ 未登录,请先执行: python oracle_skill.py login <secretKey> [clientCode]")
print("❌ 未登录,请先执行: python oracle_skill.py login <secretKey> [clientCode]", file=sys.stderr)
return False
if is_token_expired(cfg):
print("❌ 登录已过期,请重新执行: python oracle_skill.py login <secretKey> [clientCode]")
print(" 注意:同一个 secretKey 在其他设备重新登录后,本设备也需要重新登录。")
print("❌ 登录已过期,请重新执行: python oracle_skill.py login <secretKey> [clientCode]", file=sys.stderr)
print(" 注意:同一个 secretKey 在其他设备重新登录后,本设备也需要重新登录。", file=sys.stderr)
return False
return True
@@ -2403,8 +2438,10 @@ def query(server_id: str, action: str, schema: str = "", name: str = "",
Returns:
dict: {"success": True/False, "data": ..., "error": ...}
"""
if not server_id:
return {"success": False, "error": "client is required; use --client CODE"}
if not ensure_logged_in():
return {"success": False, "error": "login required"}
return {"success": False, "error": "login required", "client_code": server_id}
url = f"{get_transit_url()}/api/query"
payload = {
@@ -2424,7 +2461,7 @@ def query(server_id: str, action: str, schema: str = "", name: str = "",
logger.info(f"尝试 {attempt + 1}/{max_retries + 1}: POST {url}")
logger.debug(f"Payload: {payload}")
resp = requests.post(
resp = _http_post(
url,
json=payload,
headers=make_headers(),
@@ -2435,8 +2472,9 @@ def query(server_id: str, action: str, schema: str = "", name: str = "",
result = resp.json()
logger.info(f"请求成功: {result.get('success', 'unknown')}")
result["client_code"] = server_id
return result
except requests.exceptions.ConnectionError as e:
last_error = f"无法连接中转服务 {url},请检查地址和端口: {e}"
logger.error(last_error)
@@ -2444,14 +2482,17 @@ def query(server_id: str, action: str, schema: str = "", name: str = "",
except requests.exceptions.Timeout as e:
last_error = f"请求超时({timeout}秒),Agent 可能未响应或处理时间过长: {e}"
logger.error(last_error)
break
except requests.exceptions.HTTPError as e:
last_error = f"HTTP 错误: {e}"
logger.error(last_error)
break
except Exception as e:
last_error = f"未知错误: {e}"
logger.error(last_error)
break
# 如果不是最后一次尝试,则等待后重试
if attempt < max_retries:
@@ -2459,7 +2500,7 @@ def query(server_id: str, action: str, schema: str = "", name: str = "",
time.sleep(RETRY_DELAY)
# 所有重试都失败了
return {"success": False, "error": last_error}
return {"success": False, "error": last_error, "client_code": server_id}
def list_servers(max_retries: int = MAX_RETRIES):
@@ -2472,9 +2513,11 @@ def list_servers(max_retries: int = MAX_RETRIES):
for attempt in range(max_retries + 1):
try:
logger.info(f"获取服务器列表: {url}")
resp = requests.get(url, headers=make_headers(), timeout=10)
resp = _http_get(url, headers=make_headers(), timeout=10)
resp.raise_for_status()
return resp.json()
except (requests.exceptions.HTTPError, requests.exceptions.Timeout) as e:
return {"error": str(e)}
except Exception as e:
if attempt < max_retries:
logger.warning(f"获取服务器列表失败,重试中... ({e})")
@@ -2491,7 +2534,7 @@ def health_check(max_retries: int = MAX_RETRIES):
for attempt in range(max_retries + 1):
try:
logger.info(f"健康检查: {url}")
resp = requests.get(url, timeout=5)
resp = _http_get(url, timeout=5)
resp.raise_for_status()
return resp.json()
except Exception as e:
@@ -2549,8 +2592,9 @@ def query_with_permission(user_id: str, table_id: str, sql: str,
logger.info(f"查询用户 {user_id} 对表 {table_id} 的权限...")
# 1. 获取权限
server_id = get_server_id()
param = f"{user_id},{table_id}"
perm_result = query(get_server_id(), "get_user_perm", "BOSNDS3", param,
perm_result = query(server_id, "get_user_perm", "BOSNDS3", param,
timeout=30, max_retries=MAX_RETRIES)
if not perm_result.get("success"):
@@ -2625,13 +2669,18 @@ def query_with_permission(user_id: str, table_id: str, sql: str,
# 4. 执行最终SQL
logger.info(f"执行最终SQL: {final_sql[:100]}...")
return query(get_server_id(), "execute_query", "BOSNDS3", "",
return query(server_id, "execute_query", "BOSNDS3", "",
timeout=timeout, max_retries=MAX_RETRIES,
sql=final_sql)
def print_result(result: Dict[str, Any]):
"""格式化输出结果"""
if JSON_OUTPUT:
print(json.dumps(result, ensure_ascii=False, indent=2), file=JSON_STREAM or sys.stdout)
return
if result.get("client_code"):
print("目标客户:" + result["client_code"])
if result.get("success"):
data = result.get("data", "")
print(data)
@@ -2681,6 +2730,20 @@ def interactive_mode():
parts = line.split()
cmd = parts[0].lower()
if cmd in AGENT_COMMANDS:
# The REPL uses the same per-command target and authentication rules.
if cmd == "query" and len(parts) > 1 and parts[1].lower() in ("henlo", "renben"):
parts = [parts[0], "--client", parts[1].upper()] + parts[2:]
previous_argv = sys.argv
try:
sys.argv = [previous_argv[0]] + parts
main()
except SystemExit:
pass
finally:
sys.argv = previous_argv
continue
if cmd == "exit" or cmd == "quit":
print("再见!")
@@ -2833,8 +2896,7 @@ def interactive_mode():
if len(parts) < 2:
print(f"当前服务器: {get_server_id()}")
else:
DEFAULT_SERVER_ID = parts[1]
print(f"切换到服务器: {parts[1]}")
cmd_switch(parts[1])
elif cmd == "timeout":
if len(parts) < 2:
@@ -3036,6 +3098,7 @@ _cap_parser = argparse.ArgumentParser(prog="oracle_skill", add_help=False)
_cap_subparsers = _cap_parser.add_subparsers(dest="subcmd")
p = _cap_subparsers.add_parser('capabilities')
p.add_argument('--json', action='store_true', help='Output pure JSON')
p.add_argument('--brief', action='store_true', help='Output concise capability summary')
def cmd_capabilities(args):
"""
@@ -3542,6 +3605,16 @@ def cmd_capabilities(args):
}
}
for command in commands:
if command["name"] in AGENT_COMMANDS:
properties = command.setdefault("parameters", {"type": "object"}).setdefault("properties", {})
properties.setdefault("clientCode", {"type": "string", "description": "--client CODE before business arguments; fixes this command's target without changing the shared default."})
if command["name"] in CORE_JSON_COMMANDS:
properties.setdefault("json", {"type": "boolean", "description": "--json before business arguments; stdout contains one JSON result with client_code."})
result["global_options"] = {"client": "--client CODE (before business arguments; command-local target)", "json": "--json (before business arguments for query/metadata)"}
result["metadata"]["optional_dependencies"] = {"trusted_device_login": ["cryptography"]}
if getattr(args, "brief", False):
result["commands"] = [{"name": item["name"], "description": item["description"]} for item in commands]
# Output format
if hasattr(args, 'json') and args.json:
# Pure JSON output (no log interference)
@@ -3569,7 +3642,7 @@ def cmd_capabilities(args):
print("For machine-readable output, use: --json")
def main():
def _dispatch_main():
# Global variable declarations for main function
global DEFAULT_TIMEOUT, MAX_RETRIES, DEFAULT_SERVER_ID
@@ -3588,11 +3661,11 @@ def main():
print(json.dumps(result, indent=2, ensure_ascii=False))
elif cmd == "version":
print(f"Oracle Jump Query Skill v{VERSION}")
if len(sys.argv) >= 3 and sys.argv[2] == "agent":
args = [a for a in sys.argv[3:] if a not in ("--all", "--json")]
cmd_agent_versions(args, all_clients="--all" in sys.argv[3:], as_json="--json" in sys.argv[3:])
return
print(f"Oracle Jump Query Skill v{VERSION}")
elif cmd == "agent":
if len(sys.argv) >= 3:
@@ -3822,6 +3895,110 @@ def main():
print(__doc__)
CORE_JSON_COMMANDS = {"query", "qperm", "perm", "analyze", "source", "deps", "tables", "describe", "list", "discover", "nl2sql"}
POSITIONAL_CLIENT_INDEX = {"ops": 1, "ops_report": 0, "agent": 0,
"inspection_report": 0, "inspection_latest": 0,
"awr_status": 0, "awr_list": 0, "awr_download": 0,
"log_enable": 0, "log_disable": 0, "agent_update": 0}
AGENT_COMMANDS = CORE_JSON_COMMANDS | set(POSITIONAL_CLIENT_INDEX) | {"tablespace", "tablespaces", "log_info", "log_tail", "log_search"}
EXISTING_JSON_COMMANDS = {"inspection_report", "inspection_latest", "awr_status", "awr_list", "awr_download", "log_info", "log_tail", "log_search", "log_enable", "log_disable"}
def prepare_command(argv):
"""Parse only prefix options; never scan business arguments or SQL."""
global EXECUTION_CONTEXT, JSON_OUTPUT
EXECUTION_CONTEXT = None
JSON_OUTPUT = False
if len(argv) < 2:
return list(argv)
cmd = argv[1].lower()
if cmd not in AGENT_COMMANDS:
return list(argv)
args = list(argv[2:])
explicit = ""
as_json = False
while args and args[0] in ("--client", "--json"):
option = args.pop(0)
if option == "--client":
if not args or args[0].startswith("--"):
raise ValueError("--client 需要客户编号")
value = args.pop(0).strip()
if not value:
raise ValueError("--client 不能为空")
if explicit and explicit.upper() != value.upper():
raise ValueError("重复 --client 参数不一致")
explicit = value
else:
as_json = True
JSON_OUTPUT = as_json and cmd in CORE_JSON_COMMANDS
if as_json and cmd not in CORE_JSON_COMMANDS | EXISTING_JSON_COMMANDS:
raise ValueError("该命令不支持 --json")
required = {"query": 1, "qperm": 3, "perm": 2, "analyze": 2,
"source": 2, "deps": 2, "tables": 2, "describe": 2, "list": 1}
if len(args) < required.get(cmd, 0):
raise ValueError("业务参数不足,请参考 capabilities --json")
# Existing log parsers allow --client after the path. Retain that syntax.
if cmd in {"log_info", "log_tail", "log_search"} and "--client" in args:
index = args.index("--client")
if index + 1 >= len(args):
raise ValueError("--client 需要客户编号")
value = args[index + 1]
if explicit and explicit.upper() != value.upper():
raise ValueError("客户参数不一致")
explicit = value
position = POSITIONAL_CLIENT_INDEX.get(cmd)
positional = args[position] if position is not None and len(args) > position and not args[position].startswith("--") else ""
if cmd == "awr_download" and explicit and positional and AWR_DATE_PATTERN.fullmatch(positional):
positional = ""
if cmd == "agent_update" and explicit and positional.isdigit():
positional = ""
if explicit and positional and explicit.upper() != positional.upper():
raise ValueError("--client 与客户位置参数不一致")
cfg = get_config()
target = explicit or positional or cfg.get("client_code") or cfg.get("server_id") or DEFAULT_SERVER_ID
if not target:
raise ValueError("未指定客户,请使用 --client CODE")
target = target.upper()
EXECUTION_CONTEXT = {"client_code": target, "server_id": target,
"client_title": next((get_client_title(c) for c in cfg.get("client_list", []) if get_client_code(c).upper() == target), ""),
"transit_url": cfg.get("transit_url", TRANSIT_URL)}
JSON_OUTPUT = as_json and cmd in CORE_JSON_COMMANDS
if positional:
args[position] = target
if cmd == "awr_download" and explicit and not positional:
# Prefer the unambiguous legacy CODE DATE layout. With --client, DATE alone is accepted.
args.insert(0, target)
if cmd == "agent_update" and explicit and not positional and args:
args.insert(0, target)
if as_json and cmd not in CORE_JSON_COMMANDS:
args.append("--json")
return [argv[0], argv[1]] + args
def main():
global EXECUTION_CONTEXT, JSON_OUTPUT, JSON_STREAM
previous_argv = sys.argv
try:
sys.argv = prepare_command(sys.argv)
if JSON_OUTPUT:
# Existing progress prints stay off stdout; print_result owns JSON.
JSON_STREAM = sys.stdout
with redirect_stdout(sys.stderr):
_dispatch_main()
else:
_dispatch_main()
except (ValueError, OSError, TimeoutError) as exc:
print("❌ " + str(exc), file=sys.stderr)
if JSON_OUTPUT:
print(json.dumps({"success": False, "error": str(exc), "client_code": (EXECUTION_CONTEXT or {}).get("client_code", "")}, ensure_ascii=False))
raise SystemExit(2)
finally:
sys.argv = previous_argv
EXECUTION_CONTEXT = None
JSON_OUTPUT = False
JSON_STREAM = None
if __name__ == "__main__":
from skill_update import on_use
if on_use(SKILL_DIR):
+88
View File
@@ -0,0 +1,88 @@
"""Local configuration snapshots and bounded, process-safe updates."""
import copy
import errno
import json
import os
import tempfile
import time
from contextlib import contextmanager
class ConfigSnapshot(dict):
def __init__(self, values):
super().__init__(values)
self.original = copy.deepcopy(values)
@contextmanager
def file_lock(path, timeout=45):
# An OS lock is released even if the owning process crashes. Do not delete
# the lock file: replacing its inode would break coordination on POSIX.
with open(path, "a+b") as handle:
handle.seek(0, os.SEEK_END)
if handle.tell() == 0:
handle.write(b"\0")
handle.flush()
deadline = time.monotonic() + timeout
while True:
try:
if os.name == "nt":
import msvcrt
handle.seek(0)
msvcrt.locking(handle.fileno(), msvcrt.LK_NBLCK, 1)
else:
import fcntl
fcntl.flock(handle.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
break
except OSError as exc:
if exc.errno not in (errno.EACCES, errno.EAGAIN, errno.EDEADLK):
raise
if time.monotonic() >= deadline:
raise TimeoutError("等待本机登录/配置锁超过 %s 秒" % timeout)
time.sleep(0.05)
try:
yield
finally:
if os.name == "nt":
handle.seek(0)
msvcrt.locking(handle.fileno(), msvcrt.LK_UNLCK, 1)
else:
fcntl.flock(handle.fileno(), fcntl.LOCK_UN)
def read_config(path):
with open(path, "r", encoding="utf-8-sig") as handle:
return ConfigSnapshot(json.load(handle))
def update_config(path, values, ignored=(), initialize=False):
"""Merge only changed fields into the latest file, then atomically replace it."""
original = getattr(values, "original", {})
changes = {key: value for key, value in values.items()
if key not in ignored and (key not in original or value != original[key])}
removed = {key for key in original if key not in values and key not in ignored}
with file_lock(path + ".lock"):
try:
latest = read_config(path)
if initialize:
return latest
except FileNotFoundError:
latest = {}
if not changes and not removed and os.path.exists(path):
return ConfigSnapshot(latest)
latest.update(changes)
for key in removed:
latest.pop(key, None)
fd, temporary = tempfile.mkstemp(prefix="config.json.", suffix=".tmp", dir=os.path.dirname(path))
try:
with os.fdopen(fd, "w", encoding="utf-8") as handle:
json.dump(latest, handle, ensure_ascii=False, indent=2)
handle.flush()
os.fsync(handle.fileno())
os.replace(temporary, path)
finally:
if os.path.exists(temporary):
os.unlink(temporary)
if isinstance(values, ConfigSnapshot):
values.original = copy.deepcopy(dict(values))
return ConfigSnapshot(latest)