diff --git a/.flocks/plugins/skills/device-integration-guide/SKILL.md b/.flocks/plugins/skills/device-integration-guide/SKILL.md index f02795a9f..6ec0a2035 100644 --- a/.flocks/plugins/skills/device-integration-guide/SKILL.md +++ b/.flocks/plugins/skills/device-integration-guide/SKILL.md @@ -26,7 +26,7 @@ description: 指导 Flocks 新建、添加和接入安全设备。Use when the u ## 决策流程 -1. 已有 `device_id`:非敏感配置用 `device_manage(action="update")`,测试或排障用 `connectivity_test`;敏感字段回到设备接入页面填写。 +1. 已有 `device_id`:启停设备或更新模板声明的非密码字段时使用 `device_manage(action="update")`,测试或排障用 `connectivity_test`;密码字段回到设备接入页面填写。 2. 没有 `device_id`:先用 `device_manage(action="list")` 排除已有实例,再调用 `device_manage(action="list_templates")` 查询模板。设备实例为空不代表模板不存在。 3. 按名称、厂商、`plugin_id`、`service_id`、`storage_key` 和描述匹配模板: - `installed=true`:按 `credential_schema` 整理非敏感字段,然后进入创建流程。 @@ -58,9 +58,9 @@ description: 指导 Flocks 新建、添加和接入安全设备。Use when the u ## 页面配置与更新 -如果用户正在设备接入页面配置设备,帮助确认需要填写的表单字段,让页面负责保存。独立会话中的已有设备非敏感配置更新使用 `device_manage(action="update")`。 +如果用户正在设备接入页面配置设备,帮助确认需要填写的表单字段,让页面负责保存。独立会话中的已有设备使用 `device_manage(action="update")`:设备启停通过一级参数 `enabled` 更新,`fields` 只能包含目标模板 `credential_schema` 声明的字段。 -可以整理的非敏感字段包括: +`fields` 的具体可更新项以目标模板为准,常见字段包括: - `base_url` - `host` @@ -70,6 +70,8 @@ description: 指导 Flocks 新建、添加和接入安全设备。Use when the u - `tenant` - `region` +模板没有声明的字段不要传入 `fields`。 + 不要在聊天中索要或回显敏感字段: - `api_key` @@ -80,7 +82,7 @@ description: 指导 Flocks 新建、添加和接入安全设备。Use when the u 如果用户的目标是补填密钥、修改密码、刷新 Token 或重新登录,只说明应该在设备接入页面对应字段中处理。 -对于已有设备编辑,只处理非敏感字段和 `verify_ssl`,并说明敏感字段应在页面表单内填写。 +不要把 `enabled` 写入 `fields`。密码及 `input_type=password` 的字段不通过工具更新,应在页面表单内填写。 ## 连通性与冒烟验证 diff --git a/flocks/tool/device/manage_tool.py b/flocks/tool/device/manage_tool.py index 963debab8..aa4edf066 100644 --- a/flocks/tool/device/manage_tool.py +++ b/flocks/tool/device/manage_tool.py @@ -20,6 +20,7 @@ DeviceIntegrationCreate, DeviceIntegrationUpdate, ) +from flocks.tool.device.store import fetch_device from flocks.tool.registry import ( ParameterType, ToolCategory, @@ -39,10 +40,10 @@ "管理已接入安全设备。action=list 用于列出机房、设备、device_id 和工具集;" "action=list_templates 用于列出已有设备模板、安装状态和配置字段;" "action=create 用于从已安装模板创建设备实例,仅接受非敏感配置;" - "action=update 用于写入/更新已有设备实例的非敏感配置字段;" + "action=update 用于启停设备或更新模板声明的配置字段;" "action=connectivity_test 用于测试指定设备连通性并更新设备卡片状态。" ), - description_cn="列出设备、更新非敏感配置或测试设备连通性", + description_cn="列出设备、按模板更新设备或测试设备连通性", category=ToolCategory.SYSTEM, parameters=[ ToolParameter( @@ -50,7 +51,7 @@ type=ParameterType.STRING, description=( "操作类型:list 列出设备实例;list_templates 列出已有设备模板;" - "create 从已安装模板创建设备实例;update 更新已有设备非敏感配置;" + "create 从已安装模板创建设备实例;update 启停设备或更新模板字段;" "connectivity_test 测试设备连通性。" ), required=True, @@ -90,7 +91,7 @@ name="fields", type=ParameterType.OBJECT, description=( - "要创建或更新的非敏感设备配置字段,例如 " + "要创建或更新的模板配置字段,例如 " "{\"base_url\":\"https://device.local\"}。" "禁止传入 api_key、secret、password、token、cookie 等敏感字段。" ), @@ -102,6 +103,12 @@ description="是否开启 SSL 证书验证。action=create 或 update 时使用。", required=False, ), + ToolParameter( + name="enabled", + type=ParameterType.BOOLEAN, + description="是否启用设备。仅 action=update 时使用。", + required=False, + ), ], ) async def device_manage( @@ -113,6 +120,7 @@ async def device_manage( device_id: Optional[str] = None, fields: Optional[dict[str, Any]] = None, verify_ssl: Optional[bool] = None, + enabled: Optional[bool] = None, ) -> ToolResult: """List devices/templates, update non-secret config, or run a probe.""" normalized_action = (action or "").strip() @@ -129,7 +137,13 @@ async def device_manage( verify_ssl, ) if normalized_action == "update": - return await _update_device_config(ctx, device_id, fields, verify_ssl) + return await _update_device_config( + ctx, + device_id, + fields, + verify_ssl, + enabled, + ) if normalized_action == "connectivity_test": return await _connectivity_test(ctx, device_id) return ToolResult( @@ -301,24 +315,54 @@ async def _create_device_from_template( ) -def _normalize_update_fields( +async def _normalize_update_fields( + device_id: str, fields: Optional[dict[str, Any]], ) -> tuple[dict[str, str], Optional[str]]: if fields is None: return {}, None if not isinstance(fields, dict): return {}, "fields 必须是对象,例如 {\"base_url\":\"https://device.local\"}。" + if not fields: + return {}, None + + row = await fetch_device(device_id) + if row is None: + return {}, f"设备 {device_id!r} 未找到。" + + from flocks.tool.device.plugin_index import list_device_templates + + templates = await asyncio.to_thread(list_device_templates, refresh=False) + template = next( + ( + item + for item in templates + if item.storage_key == row["storage_key"] and item.installed + ), + None, + ) + if template is None: + return {}, "未找到该设备对应的已安装模板,无法校验 fields。" - normalized: dict[str, str] = {} - rejected: list[str] = [] - for raw_key, raw_value in fields.items(): - key = str(raw_key).strip() - if not key: - return {}, "fields 不能包含空字段名。" - if key.lower() in _SENSITIVE_FIELD_KEYS: - rejected.append(key) - continue - normalized[key] = "" if raw_value is None else str(raw_value) + schema = { + str(field.get("key") or "").strip(): field + for field in template.credential_schema + if str(field.get("key") or "").strip() + } + normalized = { + str(key).strip(): "" if value is None else str(value) + for key, value in fields.items() + } + unknown = sorted(set(normalized).difference(schema)) + if unknown: + return {}, "模板未声明字段:" + ", ".join(f"`{key}`" for key in unknown) + + rejected = sorted( + key + for key in normalized + if key.lower() in _SENSITIVE_FIELD_KEYS + or schema[key].get("input_type") == "password" + ) if rejected: return {}, ( @@ -335,6 +379,7 @@ async def _update_device_config( device_id: Optional[str], fields: Optional[dict[str, Any]], verify_ssl: Optional[bool], + enabled: Optional[bool], ) -> ToolResult: target = (device_id or "").strip() if not target: @@ -343,13 +388,13 @@ async def _update_device_config( error="action=update 时 device_id 不能为空。", ) - normalized_fields, field_error = _normalize_update_fields(fields) + normalized_fields, field_error = await _normalize_update_fields(target, fields) if field_error: return ToolResult(success=False, error=field_error) - if not normalized_fields and verify_ssl is None: + if not normalized_fields and verify_ssl is None and enabled is None: return ToolResult( success=False, - error="action=update 至少需要提供 fields 或 verify_ssl。", + error="action=update 至少需要提供 fields、verify_ssl 或 enabled。", ) log.info( @@ -359,6 +404,7 @@ async def _update_device_config( "session_id": ctx.session_id, "fields": sorted(normalized_fields), "verify_ssl": verify_ssl, + "enabled": enabled, }, ) @@ -368,6 +414,7 @@ async def _update_device_config( DeviceIntegrationUpdate( fields=normalized_fields or None, verify_ssl=verify_ssl, + enabled=enabled, ), ) except DeviceNotFoundError: @@ -401,6 +448,7 @@ async def _update_device_config( "device_id": updated.id, "updated_fields": sorted(normalized_fields), "verify_ssl": updated.verify_ssl, + "enabled": updated.enabled, }, title="设备配置已更新", ) diff --git a/flocks/updater/restart_handoff.py b/flocks/updater/restart_handoff.py index 1ed5fb050..3e1d84c32 100644 --- a/flocks/updater/restart_handoff.py +++ b/flocks/updater/restart_handoff.py @@ -19,6 +19,8 @@ DEFAULT_PORT_TIMEOUT_SECONDS = 10.0 POST_STOP_PORT_TIMEOUT_SECONDS = 20.0 SUPERVISOR_STOP_TIMEOUT_SECONDS = 20.0 +FORCED_SUPERVISOR_STOP_TIMEOUT_SECONDS = 5.0 +LEGACY_SUPERVISOR_PREPARE_TIMEOUT_SECONDS = 300.0 DEFAULT_POLL_INTERVAL_SECONDS = 0.25 @@ -102,19 +104,89 @@ def _stop_supervisor_before_restart( _record_handoff_log(f"supervisor_force_stop_failed error={terminate_exc}") return False + def stopped() -> bool: + return ( + not service_control.supervisor_is_running(paths) + and not service_manager.pid_is_running(daemon_pid) + and all(not _backend_port_in_use(port) for port in ports) + ) + deadline = time.monotonic() + timeout_seconds while time.monotonic() < deadline: - control_stopped = not service_control.supervisor_is_running(paths) - daemon_stopped = not service_manager.pid_is_running(daemon_pid) - ports_stopped = all(not _backend_port_in_use(port) for port in ports) - if control_stopped and daemon_stopped and ports_stopped: + if stopped(): return True time.sleep(poll_interval_seconds) - return ( - not service_control.supervisor_is_running(paths) - and not service_manager.pid_is_running(daemon_pid) - and all(not _backend_port_in_use(port) for port in ports) - ) + if stopped(): + return True + + if force_daemon_stop and daemon_pid is not None and service_manager.pid_is_running(daemon_pid): + try: + _record_handoff_log(f"supervisor_force_stop_after_timeout pid={daemon_pid}") + service_manager._terminate_orphan_pid(daemon_pid, "daemon", _NullConsole()) + except Exception as exc: + _record_handoff_log(f"supervisor_force_stop_failed error={exc}") + return False + + deadline = time.monotonic() + timeout_seconds + while time.monotonic() < deadline: + if stopped(): + return True + time.sleep(poll_interval_seconds) + return stopped() + + +def _prepare_legacy_supervisor_for_upgrade( + *, + daemon_pid: int | None, + service_ports: Sequence[int], + timeout_seconds: float = LEGACY_SUPERVISOR_PREPARE_TIMEOUT_SECONDS, + poll_interval_seconds: float = DEFAULT_POLL_INTERVAL_SECONDS, +) -> bool: + """Pause a v2026.7.15 supervisor after its updater parent exits.""" + from flocks.cli import service_control + + paths = service_manager.runtime_paths() + ports = set(service_ports) + + def paused() -> bool: + try: + status = service_control.read_supervisor_status(paths=paths, timeout=1.0) + except Exception: + return False + return status.backend.paused and status.webui.paused and all(not _backend_port_in_use(port) for port in ports) + + deadline = time.monotonic() + timeout_seconds + try: + response = service_control.control_api_request( + "POST", + "/upgrade/prepare", + paths=paths, + timeout=timeout_seconds, + json={}, + ) + payload = response.json() + if not isinstance(payload, dict): + raise RuntimeError("legacy supervisor returned an invalid upgrade prepare response") + status = service_control.parse_supervisor_status(payload) + if status.backend.paused and status.webui.paused and all(not _backend_port_in_use(port) for port in ports): + return True + except Exception as exc: + _record_handoff_log(f"legacy_handover_prepare_request_failed error={exc}") + status_code = getattr(getattr(exc, "response", None), "status_code", None) + if status_code in {404, 405}: + return False + if ( + not service_manager.pid_is_running(daemon_pid) + and not service_control.supervisor_is_running(paths) + and all(not _backend_port_in_use(port) for port in ports) + ): + return True + + while time.monotonic() < deadline: + if paused(): + return True + time.sleep(poll_interval_seconds) + return paused() def _parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: @@ -213,12 +285,13 @@ def _stop_services_before_upgrade(args: argparse.Namespace) -> bool: try: service_manager.stop_all(_NullConsole()) except service_manager.ServiceError as exc: - _record_handoff_log(f"service_stop_failed error={exc}") - return False + _record_handoff_log(f"service_graceful_stop_failed error={exc}") return _stop_supervisor_before_restart( daemon_pid=args.daemon_pid, backend_port=args.backend_port, service_ports=_service_ports(args), + force_daemon_stop=True, + timeout_seconds=FORCED_SUPERVISOR_STOP_TIMEOUT_SECONDS, ) @@ -277,9 +350,7 @@ def _start_service_after_upgrade(args: argparse.Namespace) -> tuple[bool, str, s stderr = updater_module._clean_process_output(completed.stderr) if completed.returncode == 0: return True, stdout, stderr - _record_handoff_log( - f"restart_failed returncode={completed.returncode} stdout={stdout} stderr={stderr}" - ) + _record_handoff_log(f"restart_failed returncode={completed.returncode} stdout={stdout} stderr={stderr}") return False, stdout, stderr @@ -591,26 +662,36 @@ def run(argv: Sequence[str] | None = None) -> int: f"frontend={args.frontend_host}:{args.frontend_port}" ) legacy_daemon_pid = _legacy_supervisor_pid(args) if args.prepare_handover else None + legacy_service_ports = tuple(sorted({args.backend_port, args.frontend_port})) supervisor_stopped = False - if args.prepare_handover: - supervisor_stopped = _stop_supervisor_before_restart( - daemon_pid=legacy_daemon_pid, - backend_port=args.backend_port, - service_ports=(args.frontend_port,), - force_daemon_stop=True, - ) - if not supervisor_stopped: - _record_handoff_log("legacy_handover_stop_timeout") - _cleanup_dir(args.cleanup_dir) - return 1 - if args.parent_pid is not None and not _wait_for_parent_exit(args.parent_pid): _record_handoff_log(f"parent_exit_timeout parent_pid={args.parent_pid}") _cleanup_dir(args.cleanup_dir) return 1 - if not args.prepare_handover and (args.pro_wheel_path or args.pro_bundle_manifest_path): + if args.prepare_handover: + legacy_prepared = _prepare_legacy_supervisor_for_upgrade( + daemon_pid=legacy_daemon_pid, + service_ports=legacy_service_ports, + timeout_seconds=max( + LEGACY_SUPERVISOR_PREPARE_TIMEOUT_SECONDS, + float(args.sync_timeout), + ), + ) + if not legacy_prepared: + _record_handoff_log("legacy_handover_pause_timeout") + supervisor_stopped = _stop_supervisor_before_restart( + daemon_pid=legacy_daemon_pid, + backend_port=args.backend_port, + service_ports=(args.frontend_port,), + force_daemon_stop=True, + ) + if not supervisor_stopped: + _record_handoff_log("legacy_handover_stop_timeout") + _cleanup_dir(args.cleanup_dir) + return 1 + elif args.pro_wheel_path or args.pro_bundle_manifest_path: supervisor_stopped = _stop_supervisor_before_restart( backend_port=args.backend_port, service_ports=(args.frontend_port,), @@ -644,7 +725,17 @@ def run(argv: Sequence[str] | None = None) -> int: _cleanup_dir(args.cleanup_dir) return 1 - if not supervisor_stopped and not _stop_supervisor_before_restart(): + if args.prepare_handover and not supervisor_stopped: + supervisor_stopped = _stop_supervisor_before_restart( + daemon_pid=legacy_daemon_pid, + backend_port=args.backend_port, + service_ports=(args.frontend_port,), + force_daemon_stop=True, + ) + elif not supervisor_stopped: + supervisor_stopped = _stop_supervisor_before_restart() + + if not supervisor_stopped: _record_handoff_log("supervisor_stop_timeout") _cleanup_dir(args.cleanup_dir) return 1 diff --git a/pyproject.toml b/pyproject.toml index f355f8608..939c3b233 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "flocks" -version = "v2026.7.22.1" +version = "v2026.7.23" description = "AI-Native SecOps platform with multi-agent collaboration" authors = [ {name = "Flocks Team", email = "team@example.com"} diff --git a/tests/skill/test_device_integration_guide_skill.py b/tests/skill/test_device_integration_guide_skill.py index e619f8462..bc9df37d3 100644 --- a/tests/skill/test_device_integration_guide_skill.py +++ b/tests/skill/test_device_integration_guide_skill.py @@ -45,3 +45,12 @@ def test_device_integration_guide_does_not_embed_page_json_protocol() -> None: assert "```json" not in content assert '"storage_key":""' not in content assert "一键回填" not in content + + +def test_device_integration_guide_updates_enabled_and_template_fields() -> None: + content = SKILL_FILE.read_text(encoding="utf-8") + + assert "设备启停通过一级参数 `enabled` 更新" in content + assert "`fields` 只能包含目标模板 `credential_schema` 声明的字段" in content + assert "不要把 `enabled` 写入 `fields`" in content + assert "密码及 `input_type=password` 的字段不通过工具更新" in content diff --git a/tests/tool/test_device_manage_tool.py b/tests/tool/test_device_manage_tool.py index 01ed78407..efd14728b 100644 --- a/tests/tool/test_device_manage_tool.py +++ b/tests/tool/test_device_manage_tool.py @@ -56,6 +56,18 @@ def make_template(**overrides) -> DeviceTemplate: return DeviceTemplate(**data) +def make_update_template() -> DeviceTemplate: + return make_template( + credential_schema=[ + {"key": "base_url", "storage": "config"}, + {"key": "username", "storage": "secret", "input_type": "text"}, + {"key": "auth_state", "storage": "config"}, + {"key": "password", "storage": "secret", "input_type": "password"}, + {"key": "api_key", "storage": "secret", "input_type": "password"}, + ] + ) + + def test_device_manage_is_registered(): tools = {tool.name for tool in ToolRegistry.list_tools()} assert "device_manage" in tools @@ -79,6 +91,7 @@ def test_device_manage_schema_includes_template_discovery_action(): "storage_key", "group_id", "fields", + "enabled", "verify_ssl", } @@ -318,20 +331,31 @@ async def test_device_manage_create_leaves_missing_fields_for_page_completion(): @pytest.mark.asyncio -async def test_device_manage_update_updates_existing_device_non_secret_config(): +async def test_device_manage_update_updates_fields_declared_by_template(): updated_device = make_device(verify_ssl=True) - with patch( - "flocks.tool.device.manage_tool.update_device", - AsyncMock(return_value=updated_device), - ) as mocked_update: + template = make_update_template() + with ( + patch( + "flocks.tool.device.manage_tool.fetch_device", + AsyncMock(return_value={"storage_key": template.storage_key}), + ), + patch( + "flocks.tool.device.plugin_index.list_device_templates", + return_value=[template], + ), + patch( + "flocks.tool.device.manage_tool.update_device", + AsyncMock(return_value=updated_device), + ) as mocked_update, + ): result = await device_manage( make_ctx(), action="update", device_id="dev-1", fields={ "base_url": "https://device.local", - "port": 443, "auth_state": "ready", + "username": "admin", }, verify_ssl=True, ) @@ -341,37 +365,108 @@ async def test_device_manage_update_updates_existing_device_non_secret_config(): assert called_device_id == "dev-1" assert update_body.fields == { "base_url": "https://device.local", - "port": "443", "auth_state": "ready", + "username": "admin", } assert update_body.verify_ssl is True assert result.success is True assert result.output["device_id"] == "dev-1" - assert result.output["updated_fields"] == ["auth_state", "base_url", "port"] + assert result.output["updated_fields"] == [ + "auth_state", + "base_url", + "username", + ] assert result.metadata["verify_ssl"] is True @pytest.mark.asyncio -async def test_device_manage_update_rejects_sensitive_fields(): +async def test_device_manage_update_changes_enabled_state(): + updated_device = make_device(enabled=False) with patch( "flocks.tool.device.manage_tool.update_device", - AsyncMock(), + AsyncMock(return_value=updated_device), ) as mocked_update: result = await device_manage( make_ctx(), action="update", device_id="dev-1", - fields={"api_key": "secret-value"}, + enabled=False, + ) + + update_body = mocked_update.await_args.args[1] + assert update_body.enabled is False + assert update_body.fields is None + assert result.success is True + assert result.output["enabled"] is False + assert result.metadata["enabled"] is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("field_name", ["password", "api_key"]) +async def test_device_manage_update_rejects_password_fields(field_name: str): + template = make_update_template() + with ( + patch( + "flocks.tool.device.manage_tool.fetch_device", + AsyncMock(return_value={"storage_key": template.storage_key}), + ), + patch( + "flocks.tool.device.plugin_index.list_device_templates", + return_value=[template], + ), + patch( + "flocks.tool.device.manage_tool.update_device", + AsyncMock(), + ) as mocked_update, + ): + result = await device_manage( + make_ctx(), + action="update", + device_id="dev-1", + fields={field_name: "secret-value"}, ) mocked_update.assert_not_awaited() assert result.success is False assert "敏感字段" in (result.error or "") - assert "api_key" in (result.error or "") + assert field_name in (result.error or "") @pytest.mark.asyncio -async def test_device_manage_update_requires_fields_or_verify_ssl(): +@pytest.mark.parametrize("field_name", ["enabled", "port"]) +async def test_device_manage_update_rejects_fields_missing_from_template( + field_name: str, +): + template = make_update_template() + with ( + patch( + "flocks.tool.device.manage_tool.fetch_device", + AsyncMock(return_value={"storage_key": template.storage_key}), + ), + patch( + "flocks.tool.device.plugin_index.list_device_templates", + return_value=[template], + ), + patch( + "flocks.tool.device.manage_tool.update_device", + AsyncMock(), + ) as mocked_update, + ): + result = await device_manage( + make_ctx(), + action="update", + device_id="dev-1", + fields={field_name: False}, + ) + + mocked_update.assert_not_awaited() + assert result.success is False + assert "模板未声明字段" in (result.error or "") + assert field_name in (result.error or "") + + +@pytest.mark.asyncio +async def test_device_manage_update_requires_fields_verify_ssl_or_enabled(): result = await device_manage( make_ctx(), action="update", @@ -379,14 +474,20 @@ async def test_device_manage_update_requires_fields_or_verify_ssl(): ) assert result.success is False - assert "至少需要提供 fields 或 verify_ssl" in (result.error or "") + assert "至少需要提供 fields、verify_ssl 或 enabled" in (result.error or "") @pytest.mark.asyncio async def test_device_manage_update_reports_missing_device_as_tool_error(): - with patch( - "flocks.tool.device.manage_tool.update_device", - AsyncMock(side_effect=DeviceNotFoundError("missing")), + with ( + patch( + "flocks.tool.device.manage_tool.fetch_device", + AsyncMock(return_value=None), + ), + patch( + "flocks.tool.device.manage_tool.update_device", + AsyncMock(), + ) as mocked_update, ): result = await device_manage( make_ctx(), @@ -395,6 +496,7 @@ async def test_device_manage_update_reports_missing_device_as_tool_error(): fields={"base_url": "https://device.local"}, ) + mocked_update.assert_not_awaited() assert result.success is False assert "未找到" in (result.error or "") diff --git a/tests/updater/test_restart_handoff.py b/tests/updater/test_restart_handoff.py index fd7058753..7f95667a1 100644 --- a/tests/updater/test_restart_handoff.py +++ b/tests/updater/test_restart_handoff.py @@ -306,6 +306,42 @@ def test_upgrade_wait_ports_exclude_legacy_cleanup_port(tmp_path: Path) -> None: assert restart_handoff._service_ports(args) == (5273,) +def test_upgrade_stop_forces_captured_daemon_after_graceful_timeout( + monkeypatch, + tmp_path: Path, +) -> None: + events: list[str] = [] + args = restart_handoff._parse_args(_simple_upgrade_handoff_args(tmp_path)) + + monkeypatch.setattr( + restart_handoff.service_manager, + "stop_all", + lambda _console: (_ for _ in ()).throw( + service_manager.ServiceError("daemon did not exit"), + ), + ) + monkeypatch.setattr( + restart_handoff, + "_stop_supervisor_before_restart", + lambda **kwargs: events.append(f"stop:{kwargs}") or True, + ) + monkeypatch.setattr( + restart_handoff, + "_record_handoff_log", + lambda message: events.append(f"log:{message}"), + ) + + assert restart_handoff._stop_services_before_upgrade(args) + assert events == [ + "log:service_graceful_stop_failed error=daemon did not exit", + ( + "stop:{'daemon_pid': 2468, 'backend_port': 5273, " + "'service_ports': (5273,), 'force_daemon_stop': True, " + "'timeout_seconds': 5.0}" + ), + ] + + def test_upgrade_install_failure_keeps_backup_and_temp_without_restart_or_rollback( monkeypatch, tmp_path: Path, @@ -674,7 +710,7 @@ def test_v2026_7_1_upgrade_handoff_runs_tasks_and_restarts(monkeypatch, tmp_path ] -def test_v2026_7_15_upgrade_handoff_stops_before_tasks_and_restarts( +def test_v2026_7_15_upgrade_handoff_pauses_before_tasks_and_restarts( monkeypatch, tmp_path: Path, ) -> None: @@ -707,8 +743,8 @@ def test_v2026_7_15_upgrade_handoff_stops_before_tasks_and_restarts( ) monkeypatch.setattr( restart_handoff, - "_ensure_backend_port_free", - lambda _backend_port: pytest.fail("legacy handover must stop the supervisor first"), + "_prepare_legacy_supervisor_for_upgrade", + lambda **kwargs: events.append(f"prepare-supervisor:{kwargs}") or True, ) monkeypatch.setattr( restart_handoff, @@ -736,14 +772,64 @@ def test_v2026_7_15_upgrade_handoff_stops_before_tasks_and_restarts( assert restart_handoff.run(args) == 0 assert events == [ + "wait-parent:1234", + ("prepare-supervisor:{'daemon_pid': 2468, 'service_ports': (5173,), 'timeout_seconds': 300.0}"), + "install", + "cleanup-handover", ( "stop-supervisor:{'daemon_pid': 2468, 'backend_port': 5173, " "'service_ports': (5173,), 'force_daemon_stop': True}" ), + f"spawn:{restart_argv}:{tmp_path}:True", + ] + + +def test_v2026_7_15_upgrade_handoff_force_stops_after_pause_failure( + monkeypatch, + tmp_path: Path, +) -> None: + events: list[str] = [] + restart_argv = ["python.exe", "-m", "flocks.cli.main", "start"] + args = _v2026_7_15_handoff_args(tmp_path, restart_argv) + + monkeypatch.setattr(restart_handoff, "_record_handoff_log", lambda message: events.append(f"log:{message}")) + monkeypatch.setattr( + restart_handoff, + "_wait_for_parent_exit", + lambda parent_pid: events.append(f"wait-parent:{parent_pid}") or True, + ) + monkeypatch.setattr( + restart_handoff, + "_prepare_legacy_supervisor_for_upgrade", + lambda **_kwargs: events.append("prepare-supervisor") or False, + ) + monkeypatch.setattr(restart_handoff, "_legacy_supervisor_pid", lambda _args: 2468) + monkeypatch.setattr( + restart_handoff, + "_stop_supervisor_before_restart", + lambda **kwargs: events.append(f"stop-supervisor:{kwargs}") or True, + ) + monkeypatch.setattr(restart_handoff, "_run_upgrade_tasks", lambda _args: events.append("install") or None) + monkeypatch.setattr(restart_handoff, "_cleanup_legacy_upgrade_handover", lambda _args: True) + monkeypatch.setattr( + restart_handoff.subprocess, + "Popen", + lambda _argv, **_kwargs: events.append("spawn") or SimpleNamespace(pid=4321), + ) + + assert restart_handoff.run(args) == 0 + assert events == [ + "log:started parent_pid=1234 backend=127.0.0.1:5173 frontend=127.0.0.1:5173", "wait-parent:1234", + "prepare-supervisor", + "log:legacy_handover_pause_timeout", + ( + "stop-supervisor:{'daemon_pid': 2468, 'backend_port': 5173, " + "'service_ports': (5173,), 'force_daemon_stop': True}" + ), "install", - "cleanup-handover", - f"spawn:{restart_argv}:{tmp_path}:True", + "spawn", + "log:restart_spawned pid=4321", ] @@ -1083,6 +1169,106 @@ def request_stop(*, paths, timeout): assert events == ["request-stop"] +def test_stop_supervisor_force_stops_after_graceful_timeout(monkeypatch) -> None: + from flocks.cli import service_control + + events: list[str] = [] + daemon_running = True + + monkeypatch.setattr(service_control, "supervisor_is_running", lambda _paths: daemon_running) + monkeypatch.setattr( + restart_handoff.service_manager, + "pid_is_running", + lambda _pid: daemon_running, + ) + monkeypatch.setattr(restart_handoff, "_backend_port_in_use", lambda _port: False) + monkeypatch.setattr( + service_control, + "request_stop", + lambda **_kwargs: events.append("request-stop"), + ) + + def terminate(pid, _label, _console) -> None: + nonlocal daemon_running + events.append(f"terminate:{pid}") + daemon_running = False + + monkeypatch.setattr(restart_handoff.service_manager, "_terminate_orphan_pid", terminate) + monkeypatch.setattr(restart_handoff, "_record_handoff_log", lambda message: events.append(f"log:{message}")) + + assert restart_handoff._stop_supervisor_before_restart( + daemon_pid=2468, + backend_port=5173, + force_daemon_stop=True, + timeout_seconds=0, + poll_interval_seconds=0, + ) + assert events == [ + "request-stop", + "log:supervisor_force_stop_after_timeout pid=2468", + "terminate:2468", + ] + + +def test_prepare_legacy_supervisor_requests_upgrade_pause(monkeypatch) -> None: + from flocks.cli import service_control + + events: list[str] = [] + paused_status = SimpleNamespace( + backend=SimpleNamespace(paused=True), + webui=SimpleNamespace(paused=True), + ) + response = SimpleNamespace(json=lambda: {"status": "paused"}) + + monkeypatch.setattr(service_control, "supervisor_is_running", lambda _paths: True) + monkeypatch.setattr(restart_handoff.service_manager, "pid_is_running", lambda _pid: True) + monkeypatch.setattr(restart_handoff, "_backend_port_in_use", lambda _port: False) + monkeypatch.setattr(service_control, "parse_supervisor_status", lambda _payload: paused_status) + + def control_api_request(method, path, **kwargs): + events.append(f"request:{method}:{path}:{kwargs['timeout']}") + return response + + monkeypatch.setattr(service_control, "control_api_request", control_api_request) + + assert restart_handoff._prepare_legacy_supervisor_for_upgrade( + daemon_pid=2468, + service_ports=(5173,), + timeout_seconds=45, + poll_interval_seconds=0, + ) + assert events == ["request:POST:/upgrade/prepare:45"] + + +def test_prepare_legacy_supervisor_polls_paused_state_after_request_timeout(monkeypatch) -> None: + from flocks.cli import service_control + + events: list[str] = [] + paused_status = SimpleNamespace( + backend=SimpleNamespace(paused=True), + webui=SimpleNamespace(paused=True), + ) + + monkeypatch.setattr(service_control, "supervisor_is_running", lambda _paths: True) + monkeypatch.setattr(restart_handoff.service_manager, "pid_is_running", lambda _pid: True) + monkeypatch.setattr(restart_handoff, "_backend_port_in_use", lambda _port: False) + monkeypatch.setattr( + service_control, + "control_api_request", + lambda *_args, **_kwargs: (_ for _ in ()).throw(TimeoutError("busy building")), + ) + monkeypatch.setattr(service_control, "read_supervisor_status", lambda **_kwargs: paused_status) + monkeypatch.setattr(restart_handoff, "_record_handoff_log", lambda message: events.append(message)) + + assert restart_handoff._prepare_legacy_supervisor_for_upgrade( + daemon_pid=2468, + service_ports=(5173,), + timeout_seconds=45, + poll_interval_seconds=0, + ) + assert events == ["legacy_handover_prepare_request_failed error=busy building"] + + @pytest.mark.skipif(sys.platform == "win32", reason="uses the Unix domain socket control API") def test_stop_supervisor_before_restart_waits_until_real_control_api_stops(monkeypatch) -> None: short_root = make_short_runtime_root("flocks-handoff-") diff --git a/uv.lock b/uv.lock index 6b9fa8690..12a019c9d 100644 --- a/uv.lock +++ b/uv.lock @@ -553,7 +553,7 @@ wheels = [ [[package]] name = "flocks" -version = "2026.7.22.1" +version = "2026.7.23" source = { editable = "." } dependencies = [ { name = "aiofiles" },