diff --git a/.gitignore b/.gitignore index 0e1aa7b..0f641b5 100644 --- a/.gitignore +++ b/.gitignore @@ -6,3 +6,7 @@ scripts/*.token scripts/update-state.json scripts/update-state.json.* scripts/update-state.lock +scripts/config.json.*.tmp +scripts/config.json.lock +scripts/config.json.auth.lock +scripts/device-key.json diff --git a/AGENTS.md b/AGENTS.md index 7583218..8daf1fd 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1,5 +1,7 @@ # oracle-jump-query Skill Release Standard +Release 1.5.54: 增加命令级客户绑定、401 一次恢复与并发续登协调;配置按字段合并并原子保存,能力支持简版、普通查询和元数据支持纯 JSON。对话只确定一次客户,由助手每次显式携带编号;保留旧用法和每日更新。仅修改 Skill,不变更 Agent、transit-server、HTTP 接口、配置 JSON 字段或审计 action、字段及保存策略,命令数保持 32。发布同步私有主库、无私有历史公共镜像与安装目录,保护本机配置和设备私钥。 + Release 1.5.53: 补充商品明细新增时 CLASSNAME 为 nds.schema.AttributeDetailSupportTableImpl 的条码输入约定,以及商品、条码、ASI 三个常用字段和 ASI 读写规则 1100000000(只在新增时可见)。本次仅更新业务文档,不修改 Agent、transit-server、CLI/HTTP 接口、配置或审计 action、字段及保存策略,命令数保持 32。发布时同步私有主库、安装目录和无私有历史公共镜像。 Release 1.5.52: 补充新建 BOS 业务单据表单的 DOCNO 生成器、STATUS/ISACTIVE 翻译器和选项组、字段读写打印规则、默认值、提交人和提交时间,以及表级 MDQSV 默认规则。ISACTIVE 默认 Y,STATUS 默认 1;仅生成方案。本次仅更新业务文档,不修改 Agent、transit-server、CLI/HTTP 接口、配置或审计 action、字段及保存策略,命令数保持 32。发布时同步私有主库、安装目录和无私有历史公共镜像。 diff --git a/CHANGELOG.md b/CHANGELOG.md index 8066af3..cd878ae 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,3 +1,11 @@ +## 1.5.54 (2026-10-08) + +- 单 Agent 命令支持 `--client CODE`,在命令开始固定目标;多步权限查询和续登重试保持同一客户,保留旧位置参数及默认客户兼容。 +- 统一 401 恢复,优先复用已续取 token,可信设备续登后只重发一次;进程间锁协调登录,配置按字段合并并原子保存,等待上限 45 秒。 +- 增加 `capabilities --brief --json` 和普通查询/元数据命令纯 JSON 输出;列表刷新和状态检查仅使用 `/api/clients`,常规查询无需前置 status/switch。 +- 查询 HTTP 错误与超时不重复执行,管理操作不增加网络失败重放;每日更新机制保留,敏感本机文件与新增锁/临时文件不提交 Git。 +- 本次仅修改 Skill CLI 参数、输出及本机运行逻辑;不修改 Agent、transit-server、HTTP 接口、配置 JSON 字段或审计 action、字段及保存策略,能力命令数保持 32;未重新生成可执行文件。 + ## 1.5.53 (2026-09-30) - 补充业务单据商品明细新增配置:`AD_TABLE.CLASSNAME='nds.schema.AttributeDetailSupportTableImpl'` 表示在商品新增输入框输入条码录入。 diff --git a/README.md b/README.md index 124a74f..75751d6 100644 --- a/README.md +++ b/README.md @@ -28,6 +28,25 @@ 开发和发布使用私有主库;每次推送私有 Git 后必须同步更新无历史公共镜像,新用户可从公共镜像下载。同步前必须检查敏感文件并完成测试。`C:\Users\qiang\.codex\skills\oracle-jump-query` 是安装后的运行目录,不作为长期源码目录。 +## 对话隔离与启动提速(1.5.54) + +每个对话只需确定一次目标客户:从用户提供的客户、工单或当前对话上下文解析出明确 client code,后续沿用,用户明确切换时更新。不得从其他对话或共享配置的“当前客户”推断本对话目标;无法确定时询问客户,名称有多个匹配时列出候选,不猜测。查询客户名称映射优先读取本地 `scripts/config.json` 的 `client_list`(只读取 code/title/name,不输出 token);找不到时执行 `clients` 刷新授权列表。 + +助手每次调用单 Agent 命令时自动在命令名后、业务参数前带 `--client CODE`,用户无需在每句话重复客户。该参数仅影响本次命令,不修改默认客户;命令开始即固定客户、中转地址,多步查询、权限获取、预检、下载和续登重试保持原目标。旧位置参数仍兼容;与 `--client` 不一致时直接报错。未指定时兼容本地默认客户,但为空则停止,不选择占位服务器或第一个客户。 + +常规查询无需先执行 `status`、`switch` 或手动登录:CLI 自动处理 token 缺失/过期,401 时复用其他进程的新 token 或可信设备续登,仅重发一次。设备未审批、撤销、过期或无本机密钥时提示原因,再由用户提供登录密钥;不自动注册设备。403、参数错误、客户离线直接返回,查询超时不重复执行。`status` 用于排障,`switch` 仅用于显式维护旧用法的默认客户。 + +普通查询及元数据命令可在业务参数前加 `--json`,stdout 只返回 JSON(含 `client_code`),登录、更新和诊断提示输出到 stderr。完整能力定义仍可用 `capabilities --json` 按需读取。 + +```powershell +python scripts/oracle_skill.py capabilities --brief --json +python scripts/oracle_skill.py describe --client WEIRUI --json BOSNDS3 M_PRODUCT +python scripts/oracle_skill.py query --client WEIRUI --json "SELECT ID, NAME FROM M_PRODUCT WHERE ID=123 AND ROWNUM<=1" +python scripts/oracle_skill.py qperm --client HENLO --json 940 12983 "SELECT ID FROM M_OTHER_INOUT WHERE ID=123 AND ROWNUM<=1" +python scripts/oracle_skill.py ops_report --client HENLO +python scripts/oracle_skill.py awr_download --client WEIRUI 20261008 +``` + ## 快速开始 更多面向日常使用的问法和操作流程,见 [操作示例](docs/操作示例.md)。 @@ -39,7 +58,7 @@ ```json { "transit_url": "https://ts.henlo.net", - "server_id": "server-001", + "server_id": "", "access_token": "", "expires_at": "" } @@ -360,7 +379,7 @@ python scripts/oracle_skill.py version agent --all --json ### 每日自动更新与使用记录 -CLI 启动时记录最后使用时间;按运行电脑本地日期,每个使用日首次启动自动 `fetch origin main`,检测到新版本时仅在 `main` 分支执行快进安装,并重新启动当前命令使用新版本。同一天后续启动不重复抓取;未使用的日期不运行后台任务。每次开始应用 Skill 时先运行 `python scripts/oracle_skill.py capabilities --json`,使只参考业务知识的任务也触发检测。 +CLI 启动时记录最后使用时间;按运行电脑本地日期,每个使用日首次启动自动 `fetch origin main`,检测到新版本时仅在 `main` 分支执行快进安装,并重新启动当前命令使用新版本。同一天后续启动不重复抓取;未使用的日期不运行后台任务。每次开始应用 Skill 时先运行 `python scripts/oracle_skill.py capabilities --brief --json`,使只参考业务知识的任务也触发检测。 本机状态文件为 `scripts/update-state.json`:`last_used_at` 记录最后使用时间,`last_checked_at` / `last_check_date` 记录检测时间和日期,`status` 记录 `checking`、`up_to_date`、`updated`、`update_blocked`、`check_failed` 或 `not_git`;检测完成时可包含本机和远端提交 ID。状态、锁和临时文件不提交 Git,不保存 token、密钥、SQL 或查询结果。 diff --git a/SKILL.md b/SKILL.md index bcf9514..f4d670b 100644 --- a/SKILL.md +++ b/SKILL.md @@ -5,13 +5,28 @@ description: Oracle 跳板查询技能。通过中转服务查询远程 Oracle # Oracle 跳板查询 -> **版本:v1.5.53** · [更新日志](./CHANGELOG.md) +> **版本:v1.5.54** · [更新日志](./CHANGELOG.md) ## 使用原则 -每次开始使用本 Skill(包括仅参考业务知识)时,先执行 `python scripts/oracle_skill.py capabilities --json`。CLI 每次启动记录最后使用时间;按运行电脑本地日期,当天第一次启动自动从当前 `origin/main` 抓取更新并尝试快进安装,成功后自动重新启动当前命令。无需登录即可检测;当天后续启动只记录使用时间。没有使用的日期不启动后台任务。 +每次开始使用本 Skill(包括仅参考业务知识)时,先执行 `python scripts/oracle_skill.py capabilities --brief --json`。CLI 每次启动记录最后使用时间;按运行电脑本地日期,当天第一次启动自动从当前 `origin/main` 抓取更新并尝试快进安装,成功后自动重新启动当前命令。无需登录即可检测;当天后续启动只记录使用时间。没有使用的日期不启动后台任务。 -使用本 skill 时,先确认当前登录状态和 client。用户指定“恒诺 / 品小二 / 未芮 / 康奈”等服务器时,可以用 `switch ` 按 code 或名称切换;找不到目标服务器时,skill 会先刷新 client 列表再重试匹配。若名称不够准确并匹配到多个 client,只列出候选项,请用户指定更准确的 client code,不要直接猜测切换。 +每个对话只需确定一次目标客户:从用户提供的客户、工单或当前对话上下文解析出明确 client code,后续沿用,用户明确切换时更新。不得从其他对话或共享配置的“当前客户”推断本对话目标;无法确定时询问客户,名称有多个匹配时列出候选,不猜测。查询客户名称映射优先读取本地 `scripts/config.json` 的 `client_list`(只读取 code/title/name,不输出 token);找不到时执行 `clients` 刷新授权列表。 + +助手每次调用单 Agent 命令时自动在命令名后、业务参数前带 `--client CODE`,用户无需在每句话重复客户。该参数仅影响本次命令,不修改默认客户;命令开始即固定客户、中转地址,多步查询、权限获取、预检、下载和续登重试保持原目标。旧位置参数仍兼容;与 `--client` 不一致时直接报错。未指定时兼容本地默认客户,但为空则停止,不选择占位服务器或第一个客户。 + +常规查询无需先执行 `status`、`switch` 或手动登录:CLI 自动处理 token 缺失/过期,401 时复用其他进程的新 token 或可信设备续登,仅重发一次。设备未审批、撤销、过期或无本机密钥时提示原因,再由用户提供登录密钥;不自动注册设备。403、参数错误、客户离线直接返回,查询超时不重复执行。`status` 用于排障,`switch` 仅用于显式维护旧用法的默认客户。 + +普通查询及元数据命令可在业务参数前加 `--json`,stdout 只返回 JSON(含 `client_code`),登录、更新和诊断提示输出到 stderr。完整能力定义仍可用 `capabilities --json` 按需读取。 + +```powershell +python scripts/oracle_skill.py capabilities --brief --json +python scripts/oracle_skill.py describe --client WEIRUI --json BOSNDS3 M_PRODUCT +python scripts/oracle_skill.py query --client WEIRUI --json "SELECT ID, NAME FROM M_PRODUCT WHERE ID=123 AND ROWNUM<=1" +python scripts/oracle_skill.py qperm --client HENLO --json 940 12983 "SELECT ID FROM M_OTHER_INOUT WHERE ID=123 AND ROWNUM<=1" +python scripts/oracle_skill.py ops_report --client HENLO +python scripts/oracle_skill.py awr_download --client WEIRUI 20261008 +``` 本 skill 的数据库通道默认只读。可以生成 SQL 和说明,但不得通过 `query` / `qperm` 执行写库语句,包括 `INSERT`、`UPDATE`、`DELETE`、`MERGE`、`DDL`、提交存储过程和其他会改变业务数据的调用。用户要求写库时,只生成 SQL,并明确说明未执行。 diff --git a/VERSION b/VERSION index a9b3eda..50ab1e3 100644 --- a/VERSION +++ b/VERSION @@ -1 +1 @@ -1.5.53 +1.5.54 diff --git a/docs/操作示例.md b/docs/操作示例.md index 3bf0a26..b8dd35f 100644 --- a/docs/操作示例.md +++ b/docs/操作示例.md @@ -11,6 +11,16 @@ - 明细表不靠猜:主表到明细表关系应通过 `AD_REFBYTABLE` 查询确认。 - 品小二中文条件:中文值直接写入 `LIKE` 可能因字符集转换返回 0 行,应改用 Oracle `UNISTR` Unicode 转义;其他 client 可按实际验证结果使用,`UNISTR` 也可作为跨环境复用时的 ASCII-safe 兜底。 +## 多对话查询 + +你可以先说“查未芮”,随后直接说“看商品表结构”“查这个过程”。助手沿用当前对话客户,并在每条命令中自动携带 `--client WEIRUI`;另一个对话查其他客户不会影响本对话。下面未带客户参数的历史示例仍兼容默认客户,助手实际调用时应补上本对话编号。 + +```powershell +python scripts/oracle_skill.py describe --client WEIRUI --json BOSNDS3 M_PRODUCT +``` + +常规查询无需先检查状态或切换共享默认客户;可信设备已审批时,失效登录由 CLI 自动恢复。 + ## 常见操作 ### 1. 检查服务状态 diff --git a/references/commands-and-auth.md b/references/commands-and-auth.md index 75c1d84..5f0e352 100644 --- a/references/commands-and-auth.md +++ b/references/commands-and-auth.md @@ -16,16 +16,16 @@ AI Skill → HTTP → Transit Server (:6357) → WebSocket → Agent → Oracle 1. 中转服务已部署并运行(默认地址: https://ts.henlo.net) 2. Agent 已部署到数据库服务器并连接中转服务 -3. 本地已安装 Python 依赖:`pip install requests` +3. 本地已安装 Python 依赖:`pip install requests`;可信设备登录另需 `pip install cryptography` ## 配置 -配置文件位于技能目录下:`D:\work\恒诺\长期支持\跳板查询\demo\ai-skill\config.json` +配置文件位于技能目录下:`scripts/config.json` ```json { "transit_url": "https://ts.henlo.net", - "server_id": "server-001", + "server_id": "", "access_token": "", "expires_at": "" } @@ -41,6 +41,37 @@ python oracle_skill.py login [clientCode] 同一个 `secretKey` 同时只允许一个设备在线。另一台设备重新登录后,当前设备的 token 会立即失效,需要重新执行 `login [clientCode]`。 +## 对话客户绑定与快速调用(1.5.54) + +每个对话只需确定一次目标客户:从用户提供的客户、工单或当前对话上下文解析出明确 client code,后续沿用,用户明确切换时更新。不得从其他对话或共享配置的“当前客户”推断本对话目标;无法确定时询问客户,名称有多个匹配时列出候选,不猜测。查询客户名称映射优先读取本地 `scripts/config.json` 的 `client_list`(只读取 code/title/name,不输出 token);找不到时执行 `clients` 刷新授权列表。 + +助手每次调用单 Agent 命令时自动在命令名后、业务参数前带 `--client CODE`,用户无需在每句话重复客户。该参数仅影响本次命令,不修改默认客户;命令开始即固定客户、中转地址,多步查询、权限获取、预检、下载和续登重试保持原目标。旧位置参数仍兼容;与 `--client` 不一致时直接报错。未指定时兼容本地默认客户,但为空则停止,不选择占位服务器或第一个客户。 + +常规查询无需先执行 `status`、`switch` 或手动登录:CLI 自动处理 token 缺失/过期,401 时复用其他进程的新 token 或可信设备续登,仅重发一次。设备未审批、撤销、过期或无本机密钥时提示原因,再由用户提供登录密钥;不自动注册设备。403、参数错误、客户离线直接返回,查询超时不重复执行。`status` 用于排障,`switch` 仅用于显式维护旧用法的默认客户。 + +普通查询及元数据命令可在业务参数前加 `--json`,stdout 只返回 JSON(含 `client_code`),登录、更新和诊断提示输出到 stderr。完整能力定义仍可用 `capabilities --json` 按需读取。 + +```powershell +python scripts/oracle_skill.py capabilities --brief --json +python scripts/oracle_skill.py describe --client WEIRUI --json BOSNDS3 M_PRODUCT +python scripts/oracle_skill.py query --client WEIRUI --json "SELECT ID, NAME FROM M_PRODUCT WHERE ID=123 AND ROWNUM<=1" +python scripts/oracle_skill.py qperm --client HENLO --json 940 12983 "SELECT ID FROM M_OTHER_INOUT WHERE ID=123 AND ROWNUM<=1" +python scripts/oracle_skill.py ops_report --client HENLO +python scripts/oracle_skill.py awr_download --client WEIRUI 20261008 +``` + +本机 token 仍由同一安装目录中的 `scripts/config.json` 共享;不同对话共享身份、各自绑定目标。续登锁和配置锁最多等待 45 秒,锁由操作系统在进程退出时释放。配置按最新文件合并修改字段,并用同目录临时文件原子替换;锁、临时文件及设备私钥不提交 Git。不同安装目录间的认证共享未在本次实现。`logout` 会退出共享同一安装配置的所有对话。 + +`--client` 适用于 query/qperm/perm、analyze/source/deps/tables/describe/list/discover/nl2sql、agent、ops/ops_report、tablespace(s)、inspection_report/inspection_latest、awr_status/awr_list/awr_download、日志命令及 agent_update。历史报告 `inspection_get` 仍按报告 ID 和服务端授权访问;批量 Agent 版本仍按逐客户参数或 `--all` 访问。 + +`--json` 的新增前缀用法适用于 query/qperm/perm、analyze/source/deps/tables/describe/list/discover/nl2sql;巡检、AWR 和日志命令保留已有 JSON 用法。SQL 内的 `--client`、`--json` 文本不会被当作选项。 + +### 旧用法兼容说明 + +使用本 skill 时,先确认当前登录状态和 client。用户指定“恒诺 / 品小二 / 未芮 / 康奈”等服务器时,可以用 `switch ` 按 code 或名称切换;找不到目标服务器时,skill 会先刷新 client 列表再重试匹配。若名称不够准确并匹配到多个 client,只列出候选项,请用户指定更准确的 client code,不要直接猜测切换。 + +上述先检查、再切换的旧操作仍可用于排障及维护默认客户;多对话查询使用本节的显式目标流程。手动 `device_login [clientCode]` 可强制重新申请 token,并协调同时启动的设备登录;普通命令自动续登无需手动调用。手动密钥登录保留原有选择默认客户的语义。 + ## 调用方式 使用 `scripts/oracle_skill.py` 脚本: @@ -108,7 +139,7 @@ python scripts/oracle_skill.py curl -X POST https://ts.henlo.net/api/query \ -H "Content-Type: application/json" \ -d '{ - "server_id": "server-001", + "server_id": "", "action": "analyze_procedure", "schema": "BOS", "name": "M_RETAIL_SUBMIT", @@ -122,7 +153,7 @@ curl -X POST https://ts.henlo.net/api/query \ curl -X POST https://ts.henlo.net/api/query \ -H "Content-Type: application/json" \ -d '{ - "server_id": "server-001", + "server_id": "", "action": "describe_table", "schema": "bosnds3", "name": "xcx_so", diff --git a/scripts/oracle_skill.py b/scripts/oracle_skill.py index 5b989c0..4a3d19b 100644 --- a/scripts/oracle_skill.py +++ b/scripts/oracle_skill.py @@ -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 [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 [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 [clientCode]") + print("❌ 未登录,请先执行: python oracle_skill.py login [clientCode]", file=sys.stderr) return False if is_token_expired(cfg): - print("❌ 登录已过期,请重新执行: python oracle_skill.py login [clientCode]") - print(" 注意:同一个 secretKey 在其他设备重新登录后,本设备也需要重新登录。") + print("❌ 登录已过期,请重新执行: python oracle_skill.py login [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): diff --git a/scripts/runtime_state.py b/scripts/runtime_state.py new file mode 100644 index 0000000..dc6251f --- /dev/null +++ b/scripts/runtime_state.py @@ -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) diff --git a/tests/test_runtime.py b/tests/test_runtime.py new file mode 100644 index 0000000..1b39e3f --- /dev/null +++ b/tests/test_runtime.py @@ -0,0 +1,390 @@ +import io +import base64 +import json +import os +from pathlib import Path +import subprocess +import sys +import tempfile +import time +import unittest +from contextlib import redirect_stdout, redirect_stderr +from unittest import mock + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from scripts import oracle_skill as skill +from scripts.runtime_state import file_lock, read_config, update_config + + +def response(status=200, payload=None): + result = mock.Mock(status_code=status, ok=status < 400) + result.json.return_value = payload or {"success": True, "data": "ok"} + if status >= 400: + result.raise_for_status.side_effect = skill.requests.HTTPError(str(status)) + return result + + +def worker(mode, directory, target): + directory = Path(directory) + skill.get_script_dir = lambda: str(directory) + cfg = skill.get_config() + (directory / (target + '.ready')).touch() + deadline = time.monotonic() + 15 + while not (directory / 'go').exists(): + if time.monotonic() >= deadline: + raise RuntimeError('worker barrier timeout') + time.sleep(.01) + if mode == 'renew': + skill.prepare_command(['skill', 'query', '--client', target, 'SELECT 1 FROM DUAL']) + def renew(values): + with open(directory / 'renewals', 'a') as handle: + handle.write('renew\n') + time.sleep(.2) + values.update(access_token='new-token', expires_at='2999-01-01T00:00:00+00:00') + skill.save_config(values) + return True + skill._auto_device_login_unlocked = renew + rejected = 'old-token' if cfg.get('expires_at', '').startswith('2999') else None + assert skill._auto_device_login(skill.get_config(), rejected) + def post(url, **kwargs): + assert kwargs['headers']['Authorization'] == 'Bearer new-token' + assert kwargs['json']['server_id'] == target + return response() + skill.requests.post = post + assert skill.query(skill.get_server_id(), 'execute_query', sql='SELECT 1 FROM DUAL')['client_code'] == target + elif mode == 'merge': + cfg[target] = target + skill.save_config(cfg) + elif mode == 'crash': + with file_lock(str(directory / 'crash.lock')): + (directory / 'locked').touch() + time.sleep(20) + + +class RuntimeTests(unittest.TestCase): + def setUp(self): + self.temp = tempfile.TemporaryDirectory() + self.addCleanup(self.temp.cleanup) + self.directory = Path(self.temp.name) + self.path = self.directory / 'config.json' + self.path.write_text(json.dumps({ + 'transit_url': 'https://transit.example', 'access_token': 'old-token', + 'expires_at': '2999-01-01T00:00:00+00:00', + 'client_code': 'HENLO', 'server_id': 'HENLO', 'user_name': 'test', + 'client_list': [{'code': 'HENLO'}, {'code': 'WEIRUI'}], + })) + patch = mock.patch.object(skill, 'get_script_dir', return_value=str(self.directory)) + patch.start() + self.addCleanup(patch.stop) + skill.EXECUTION_CONTEXT = None + skill.JSON_OUTPUT = False + self.addCleanup(self.reset_context) + + def reset_context(self): + skill.EXECUTION_CONTEXT = None + skill.JSON_OUTPUT = False + + def prepare(self, *args): + return skill.prepare_command(['skill', *args]) + + def test_explicit_target_does_not_persist_and_survives_default_switch(self): + self.prepare('query', '--client', 'WEIRUI', 'SELECT 1 FROM DUAL') + update_config(str(self.path), {'client_code': 'OTHER', 'server_id': 'OTHER'}) + cfg = skill.get_config() + cfg['access_token'] = 'new-token' + skill.save_config(cfg) + self.assertEqual(skill.get_server_id(), 'WEIRUI') + stored = read_config(str(self.path)) + self.assertEqual(stored['client_code'], 'OTHER') + self.assertEqual(stored['access_token'], 'new-token') + + def test_implicit_default_is_fixed_at_command_start(self): + self.prepare('describe', 'BOSNDS3', 'M_PRODUCT') + update_config(str(self.path), {'client_code': 'WEIRUI', 'server_id': 'WEIRUI', 'transit_url': 'https://other.example'}) + self.assertEqual(skill.get_server_id(), 'HENLO') + self.assertEqual(skill.get_transit_url(), 'https://transit.example') + + def test_missing_target_stops_before_login(self): + update_config(str(self.path), {'client_code': '', 'server_id': ''}) + with mock.patch.object(skill.requests, 'post') as post: + with self.assertRaisesRegex(ValueError, '客户'): + self.prepare('query', 'SELECT 1 FROM DUAL') + post.assert_not_called() + + def test_json_missing_target_returns_structured_error(self): + update_config(str(self.path), {'client_code': '', 'server_id': ''}) + output = io.StringIO() + with mock.patch.object(sys, 'argv', ['skill', 'query', '--json', 'SELECT 1 FROM DUAL']), redirect_stdout(output), redirect_stderr(io.StringIO()), self.assertRaises(SystemExit): + skill.main() + self.assertFalse(json.loads(output.getvalue())['success']) + + def test_empty_explicit_target_and_missing_arguments_are_rejected(self): + for args in [('query', '--client', '', 'SELECT 1 FROM DUAL'), ('describe', '--client', 'HENLO', 'BOSNDS3')]: + with self.subTest(args=args), self.assertRaises(ValueError): + self.prepare(*args) + + def test_positional_conflict_and_legacy_layouts(self): + with self.assertRaisesRegex(ValueError, '不一致'): + self.prepare('ops_report', '--client', 'HENLO', 'WEIRUI') + self.assertEqual(self.prepare('ops_report', 'weirui')[-1], 'WEIRUI') + self.assertEqual(self.prepare('awr_download', '--client', 'HENLO', '20261008')[-2:], ['HENLO', '20261008']) + self.assertEqual(self.prepare('agent_update', '--client', 'HENLO', '300')[-2:], ['HENLO', '300']) + + def test_sql_is_not_scanned_for_options(self): + sql = "SELECT '--client WEIRUI --json' FROM DUAL" + argv = self.prepare('query', '--client', 'HENLO', sql) + self.assertEqual(argv[-1], sql) + self.assertFalse(skill.JSON_OUTPUT) + + def test_permission_lookup_keeps_customer_between_steps(self): + calls = [] + def query(server, action, *args, **kwargs): + calls.append(server) + update_config(str(self.path), {'client_code': 'WEIRUI', 'server_id': 'WEIRUI'}) + return {'success': True, 'data': '{}'} + with mock.patch.object(skill, 'query', side_effect=query): + skill.query_with_permission('1', '2', 'SELECT 1 FROM DUAL') + self.assertEqual(calls, ['HENLO', 'HENLO']) + + def test_valid_token_requires_only_business_request(self): + self.prepare('query', '--client', 'HENLO', 'SELECT 1 FROM DUAL') + with mock.patch.object(skill.requests, 'get') as get, mock.patch.object(skill.requests, 'post', return_value=response()) as post: + self.assertTrue(skill.query('HENLO', 'execute_query', sql='SELECT 1 FROM DUAL')['success']) + get.assert_not_called() + self.assertEqual(post.call_count, 1) + + def test_401_recovers_once_and_preserves_payload(self): + self.prepare('query', '--client', 'WEIRUI', 'SELECT 1 FROM DUAL') + def renew(cfg): + cfg['access_token'] = 'new-token' + skill.save_config(cfg) + return True + with mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew) as login, mock.patch.object(skill.requests, 'post', side_effect=[response(401), response()]) as post: + result = skill.query('WEIRUI', 'execute_query', sql='SELECT 1 FROM DUAL') + self.assertTrue(result['success']) + self.assertEqual(login.call_count, 1) + self.assertEqual(post.call_count, 2) + self.assertEqual(post.call_args.kwargs['headers']['Authorization'], 'Bearer new-token') + self.assertEqual(post.call_args_list[0].kwargs['json'], post.call_args_list[1].kwargs['json']) + + def test_second_401_and_403_do_not_sleep_or_loop(self): + for statuses in ([401, 401], [403]): + with self.subTest(statuses=statuses): + self.reset_context() + self.prepare('query', '--client', 'HENLO', 'SELECT 1 FROM DUAL') + def renew(cfg): + cfg['access_token'] += '-new' + skill.save_config(cfg) + return True + with mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew), mock.patch.object(skill.requests, 'post', side_effect=[response(status) for status in statuses]) as post, mock.patch.object(skill.time, 'sleep') as sleep: + self.assertFalse(skill.query('HENLO', 'execute_query')['success']) + self.assertEqual(post.call_count, len(statuses)) + sleep.assert_not_called() + + def test_rejected_old_token_reuses_already_renewed_token(self): + update_config(str(self.path), {'access_token': 'new-token'}) + with mock.patch.object(skill, '_auto_device_login_unlocked') as login: + cfg = skill.get_config() + self.assertTrue(skill._auto_device_login(cfg, rejected_token='old-token')) + login.assert_not_called() + self.assertEqual(cfg['access_token'], 'new-token') + + def test_revoked_device_returns_failure(self): + with mock.patch.object(skill, '_auto_device_login_unlocked', return_value=False), mock.patch.object(skill.requests, 'post', return_value=response(401)) as post: + self.assertFalse(skill.query('HENLO', 'execute_query')['success']) + self.assertEqual(post.call_count, 1) + + def test_real_device_challenge_signing_preserves_default_customer(self): + from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey + from cryptography.hazmat.primitives.serialization import Encoding, PrivateFormat, NoEncryption + private = Ed25519PrivateKey.generate() + key = {'device_id': 'test-device', 'key_fingerprint': 'test-fingerprint', + 'private_key': base64.b64encode(private.private_bytes(Encoding.Raw, PrivateFormat.Raw, NoEncryption())).decode('ascii')} + self.prepare('query', '--client', 'WEIRUI', 'SELECT 1 FROM DUAL') + update_config(str(self.path), {'expires_at': '2000-01-01T00:00:00+00:00'}) + responses = [response(payload={'success': True, 'challenge': 'test-challenge'}), + response(payload={'success': True, 'access_token': 'signed-token', 'expires_at': '2999-01-01T00:00:00+00:00', 'current_client': {'code': 'WEIRUI'}})] + with mock.patch.object(skill, '_load_device_key', return_value=key), mock.patch.object(skill.requests, 'post', side_effect=responses) as post: + self.assertTrue(skill.ensure_logged_in()) + login = post.call_args.kwargs['json'] + private.public_key().verify(base64.b64decode(login['signature']), b'test-challenge') + self.assertEqual(login['client_code'], 'WEIRUI') + cfg = read_config(str(self.path)) + self.assertEqual(cfg['access_token'], 'signed-token') + self.assertEqual(cfg['client_code'], 'HENLO') + + def test_device_denial_is_diagnostic_and_does_not_fallback(self): + from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey + from cryptography.hazmat.primitives.serialization import Encoding, PrivateFormat, NoEncryption + private = Ed25519PrivateKey.generate() + key = {'device_id': 'test-device', 'key_fingerprint': 'test-fingerprint', + 'private_key': base64.b64encode(private.private_bytes(Encoding.Raw, PrivateFormat.Raw, NoEncryption())).decode('ascii')} + cfg = skill.get_config() + stderr = io.StringIO() + with mock.patch.object(skill, '_load_device_key', return_value=key), mock.patch.object(skill.requests, 'post', return_value=response(403, {'success': False, 'error': 'device revoked'})) as post, redirect_stderr(stderr): + self.assertFalse(skill._auto_device_login(cfg, rejected_token='old-token')) + self.assertEqual(post.call_count, 1) + self.assertIn('device revoked', stderr.getvalue()) + self.assertEqual(cfg['access_token'], 'old-token') + + def test_management_transport_failure_is_not_replayed(self): + cfg = skill.get_config() + cfg['upgrade_admin_token'] = 'test-admin-token' + skill.save_config(cfg) + with mock.patch.object(skill.requests, 'post', side_effect=skill.requests.Timeout) as post: + self.assertFalse(skill.agent_update('HENLO')['success']) + self.assertEqual(post.call_count, 1) + + def test_admin_401_does_not_renew_user_session(self): + cfg = skill.get_config() + cfg['upgrade_admin_token'] = 'test-admin-token' + skill.save_config(cfg) + with mock.patch.object(skill.requests, 'post', return_value=response(401)) as post, mock.patch.object(skill, '_auto_device_login') as login: + self.assertFalse(skill.agent_update('HENLO')['success']) + self.assertEqual(post.call_count, 1) + login.assert_not_called() + + def test_timeout_not_replayed_and_connection_failures_are_bounded(self): + with mock.patch.object(skill.requests, 'post', side_effect=skill.requests.Timeout) as post, mock.patch.object(skill.time, 'sleep') as sleep: + self.assertFalse(skill.query('HENLO', 'execute_query')['success']) + self.assertEqual(post.call_count, 1) + sleep.assert_not_called() + with mock.patch.object(skill.requests, 'post', side_effect=skill.requests.ConnectionError) as post, mock.patch.object(skill.time, 'sleep'): + self.assertFalse(skill.query('HENLO', 'execute_query', max_retries=1)['success']) + self.assertEqual(post.call_count, 2) + + def test_json_query_with_renewal_has_clean_stdout(self): + stdout, stderr = io.StringIO(), io.StringIO() + def renew(cfg): + print('renew progress', file=sys.stderr) + cfg['access_token'] = 'new-token' + skill.save_config(cfg) + return True + argv = ['skill', 'query', '--client', 'WEIRUI', '--json', 'SELECT 1 FROM DUAL'] + with mock.patch.object(sys, 'argv', argv), mock.patch.object(skill.requests, 'post', side_effect=[response(401), response()]), mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew), redirect_stdout(stdout), redirect_stderr(stderr): + skill.main() + self.assertEqual(json.loads(stdout.getvalue())['client_code'], 'WEIRUI') + self.assertIn('renew progress', stderr.getvalue()) + + def test_status_and_refresh_use_one_clients_request(self): + with mock.patch.object(skill.requests, 'get', return_value=response(payload={'success': True, 'clients': [{'code': 'HENLO', 'online': True}]})) as get, redirect_stdout(io.StringIO()): + skill.cmd_status() + self.assertEqual(get.call_count, 1) + self.assertTrue(get.call_args.args[0].endswith('/api/clients')) + + def test_brief_capabilities_preserve_full_contract(self): + results = [] + for brief in (False, True): + output = io.StringIO() + with redirect_stdout(output): + skill.cmd_capabilities(type('Args', (), {'json': True, 'brief': brief})()) + results.append(output.getvalue()) + full, brief = map(json.loads, results) + self.assertEqual([c['name'] for c in full['commands']], [c['name'] for c in brief['commands']]) + self.assertIn('parameters', full['commands'][0]) + self.assertNotIn('parameters', brief['commands'][0]) + self.assertLess(len(results[1]), len(results[0]) // 2) + + def test_stale_snapshots_merge_only_changed_fields(self): + a, b = skill.get_config(), skill.get_config() + a['access_token'] = 'new-token' + b['client_code'] = b['server_id'] = 'WEIRUI' + skill.save_config(a) + skill.save_config(b) + result = read_config(str(self.path)) + self.assertEqual(result['access_token'], 'new-token') + self.assertEqual(result['client_code'], 'WEIRUI') + + def test_initialization_does_not_replace_existing_credentials(self): + cfg = update_config(str(self.path), {'access_token': '', 'client_code': ''}, initialize=True) + self.assertEqual(cfg['access_token'], 'old-token') + self.assertEqual(read_config(str(self.path))['client_code'], 'HENLO') + + def test_template_initialization_keeps_all_fields(self): + self.path.unlink() + (self.directory / 'config.template.json').write_text(json.dumps({'transit_url': 'https://template.example', 'server_id': '', 'custom': 'keep'})) + with mock.patch.object(skill, 'read_config', wraps=read_config) as read: + # The CWD fallback may exist in the repository; make only that path absent. + def read_without_cwd(path): + if path == 'config.json': + raise FileNotFoundError(path) + return read_config(path) + read.side_effect = read_without_cwd + cfg = skill.get_config() + self.assertEqual(cfg['custom'], 'keep') + self.assertEqual(read_config(str(self.path))['transit_url'], 'https://template.example') + + def test_streaming_get_401_keeps_download_options(self): + old, new = response(401), response() + def renew(cfg): + cfg['access_token'] = 'new-token' + skill.save_config(cfg) + return True + with mock.patch.object(skill.requests, 'get', side_effect=[old, new]) as get, mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew): + result = skill._http_get('https://transit.example/api/awr/download', params={'client_code': 'WEIRUI', 'date': '20261008'}, headers=skill.make_headers(), stream=True, timeout=120) + self.assertIs(result, new) + old.close.assert_called_once() + self.assertTrue(get.call_args.kwargs['stream']) + self.assertEqual(get.call_args.kwargs['params']['client_code'], 'WEIRUI') + + def test_lock_wait_is_bounded(self): + path = str(self.directory / 'test.lock') + with file_lock(path): + with self.assertRaises(TimeoutError): + with file_lock(path, timeout=.1): + self.fail('lock acquired twice') + + def start_workers(self, mode): + processes = [subprocess.Popen([sys.executable, __file__, '--worker', mode, str(self.directory), target], stdout=subprocess.PIPE, stderr=subprocess.PIPE) for target in ('HENLO', 'WEIRUI')] + for process in processes: + self.addCleanup(lambda p=process: p.kill() if p.poll() is None else None) + deadline = time.monotonic() + 15 + while not all((self.directory / (target + '.ready')).exists() for target in ('HENLO', 'WEIRUI')): + if time.monotonic() > deadline or any(p.poll() is not None for p in processes): + self.fail('workers failed to start') + time.sleep(.01) + (self.directory / 'go').touch() + return processes + + def assert_workers(self, processes): + for process in processes: + stdout, stderr = process.communicate(timeout=15) + self.assertEqual(process.returncode, 0, stderr.decode('utf-8', errors='replace')) + + def test_two_processes_share_one_expired_login_and_keep_targets(self): + update_config(str(self.path), {'expires_at': '2000-01-01T00:00:00+00:00'}) + processes = self.start_workers('renew') + update_config(str(self.path), {'client_code': 'OTHER', 'server_id': 'OTHER'}) + self.assert_workers(processes) + self.assertEqual((self.directory / 'renewals').read_text().splitlines(), ['renew']) + self.assertEqual(read_config(str(self.path))['client_code'], 'OTHER') + + def test_two_processes_share_one_401_renewal(self): + self.assert_workers(self.start_workers('renew')) + self.assertEqual((self.directory / 'renewals').read_text().splitlines(), ['renew']) + + def test_two_processes_merge_configuration_updates(self): + self.assert_workers(self.start_workers('merge')) + cfg = read_config(str(self.path)) + self.assertEqual(cfg['HENLO'], 'HENLO') + self.assertEqual(cfg['WEIRUI'], 'WEIRUI') + + def test_crashed_process_releases_lock(self): + processes = self.start_workers('crash') + deadline = time.monotonic() + 10 + while not (self.directory / 'locked').exists(): + if time.monotonic() > deadline: + self.fail('lock not acquired') + time.sleep(.01) + for process in processes: + process.kill() + process.communicate(timeout=5) + with file_lock(str(self.directory / 'crash.lock'), timeout=.5): + pass + + +if __name__ == '__main__': + if len(sys.argv) > 1 and sys.argv[1] == '--worker': + worker(*sys.argv[2:]) + else: + unittest.main() diff --git a/tests/test_skill_update.py b/tests/test_skill_update.py index 6e0cfb3..536224b 100644 --- a/tests/test_skill_update.py +++ b/tests/test_skill_update.py @@ -120,7 +120,7 @@ class DailyUpdateTests(unittest.TestCase): source, installed = self.repository() (source / "scripts").mkdir() scripts = pathlib.Path(skill_update.__file__).parent - for name in ("oracle_skill.py", "skill_update.py"): + for name in ("oracle_skill.py", "skill_update.py", "runtime_state.py"): shutil.copyfile(str(scripts / name), str(source / "scripts" / name)) self.commit(source, "add cli") self.git(source, "push", "origin", "main")