同步客户绑定与登录恢复优化
This commit is contained in:
+317
-140
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user