From a7d2677bc001c2f2dec79afbb1dae29e3062ca8e Mon Sep 17 00:00:00 2001 From: Petr Date: Sat, 28 Feb 2026 01:31:08 +0100 Subject: [PATCH] Phase 4: Architecture cleanup - extract DoctorService, fix gather, refactor lineage - Extract service-layer logic from commands/doctor.py into DoctorService (17 lines vs 285) - Move doctor formatting to output.py (format_doctor_panel) for consistency - Split redundant except (KeboolaApiError, Exception) into separate blocks in org_service - Replace sequential task collection in MCP service with asyncio.gather + shared _gather_results helper - Refactor LineageService to use BaseService._run_parallel() instead of custom ThreadPoolExecutor - Wire DoctorService into cli.py context - Add comprehensive DoctorService unit tests (test_doctor_service.py) --- src/keboola_agent_cli/cli.py | 3 + src/keboola_agent_cli/commands/doctor.py | 279 +----------- src/keboola_agent_cli/output.py | 36 ++ .../services/doctor_service.py | 253 +++++++++++ .../services/lineage_service.py | 64 +-- src/keboola_agent_cli/services/mcp_service.py | 76 ++-- src/keboola_agent_cli/services/org_service.py | 9 +- tests/test_cli.py | 4 +- tests/test_doctor_service.py | 397 ++++++++++++++++++ 9 files changed, 762 insertions(+), 359 deletions(-) create mode 100644 src/keboola_agent_cli/services/doctor_service.py create mode 100644 tests/test_doctor_service.py diff --git a/src/keboola_agent_cli/cli.py b/src/keboola_agent_cli/cli.py index 729051d2..528e8fbe 100644 --- a/src/keboola_agent_cli/cli.py +++ b/src/keboola_agent_cli/cli.py @@ -15,6 +15,7 @@ from .config_store import ConfigStore from .output import OutputFormatter from .services.config_service import ConfigService +from .services.doctor_service import DoctorService from .services.job_service import JobService from .services.lineage_service import LineageService from .services.mcp_service import McpService @@ -76,6 +77,7 @@ def main( lineage_service = LineageService(config_store=config_store) org_service = OrgService(config_store=config_store) mcp_service = McpService(config_store=config_store) + doctor_service = DoctorService(config_store=config_store, mcp_service=mcp_service) ctx.ensure_object(dict) ctx.obj["formatter"] = formatter @@ -89,3 +91,4 @@ def main( ctx.obj["lineage_service"] = lineage_service ctx.obj["org_service"] = org_service ctx.obj["mcp_service"] = mcp_service + ctx.obj["doctor_service"] = doctor_service diff --git a/src/keboola_agent_cli/commands/doctor.py b/src/keboola_agent_cli/commands/doctor.py index c11398ec..9e5c4bbe 100644 --- a/src/keboola_agent_cli/commands/doctor.py +++ b/src/keboola_agent_cli/commands/doctor.py @@ -1,284 +1,17 @@ -"""Doctor command - comprehensive health check for CLI configuration and connectivity. +"""Doctor command - thin CLI wrapper over DoctorService. -Runs four checks: -1. Config file existence and permissions (0600) -2. Config file valid JSON and parseable -3. Token verification for each project (API call with response time) -4. CLI version +Delegates all health check logic to DoctorService (3-layer architecture). """ -import json -import os -import stat -import time -from typing import Any - import typer -from rich.console import Console -from rich.panel import Panel -from .. import __version__ -from ..config_store import ConfigStore -from ..errors import KeboolaApiError -from ..models import AppConfig -from ..services.mcp_service import McpService -from ..services.base import ClientFactory, default_client_factory +from ..output import format_doctor_panel from ._helpers import get_formatter, get_service -def _check_config_file(config_store: ConfigStore) -> dict[str, Any]: - """Check 1: Config file exists and has correct permissions (0600). - - Returns: - Dict with check name, status (pass/fail/warn), and message. - """ - config_path = config_store.config_path - - if not config_path.exists(): - return { - "check": "config_file", - "name": "Config file", - "status": "warn", - "message": f"Config file not found at {config_path}. Run 'kbagent project add' to create it.", - } - - # Check permissions (Unix only) - try: - file_stat = os.stat(config_path) - mode = stat.S_IMODE(file_stat.st_mode) - if mode != 0o600: - return { - "check": "config_file", - "name": "Config file", - "status": "warn", - "message": f"Config file exists at {config_path} but has permissions {oct(mode)} (expected 0o600).", - } - except OSError: - # On platforms where permission checking is not reliable - pass - - return { - "check": "config_file", - "name": "Config file", - "status": "pass", - "message": f"Config file exists at {config_path} with correct permissions.", - } - - -def _check_config_valid(config_store: ConfigStore) -> tuple[dict[str, Any], AppConfig | None]: - """Check 2: Config file is valid JSON and parseable. - - Returns: - Tuple of (check result dict, parsed AppConfig or None on failure). - """ - config_path = config_store.config_path - - if not config_path.exists(): - return { - "check": "config_valid", - "name": "Config parseable", - "status": "skip", - "message": "No config file to validate.", - }, None - - try: - raw = config_path.read_text(encoding="utf-8") - except OSError as exc: - return { - "check": "config_valid", - "name": "Config parseable", - "status": "fail", - "message": f"Cannot read config file: {exc}", - }, None - - try: - json.loads(raw) - except json.JSONDecodeError as exc: - return { - "check": "config_valid", - "name": "Config parseable", - "status": "fail", - "message": f"Config file is not valid JSON: {exc}", - }, None - - try: - config = config_store.load() - except Exception as exc: - return { - "check": "config_valid", - "name": "Config parseable", - "status": "fail", - "message": f"Config file has invalid structure: {exc}", - }, None - - project_count = len(config.projects) - return { - "check": "config_valid", - "name": "Config parseable", - "status": "pass", - "message": f"Config file is valid JSON with {project_count} project(s).", - }, config - - -def _check_connectivity( - config: AppConfig | None, - client_factory: ClientFactory, -) -> list[dict[str, Any]]: - """Check 3: For each project, verify token via API call. - - Returns: - List of check result dicts, one per project. - """ - if config is None or not config.projects: - return [ - { - "check": "connectivity", - "name": "Project connectivity", - "status": "skip", - "message": "No projects configured.", - } - ] - - results = [] - for alias, project in config.projects.items(): - client = client_factory(project.stack_url, project.token) - start_time = time.monotonic() - try: - token_info = client.verify_token() - elapsed = time.monotonic() - start_time - results.append( - { - "check": "connectivity", - "name": f"Project '{alias}'", - "status": "pass", - "message": ( - f"Connected to {project.stack_url} " - f"(project: {token_info.project_name}, id: {token_info.project_id}) " - f"in {round(elapsed * 1000)}ms" - ), - "alias": alias, - "response_time_ms": round(elapsed * 1000), - } - ) - except KeboolaApiError as exc: - elapsed = time.monotonic() - start_time - results.append( - { - "check": "connectivity", - "name": f"Project '{alias}'", - "status": "fail", - "message": f"Failed: {exc.message}", - "alias": alias, - "error_code": exc.error_code, - "response_time_ms": round(elapsed * 1000), - } - ) - finally: - client.close() - - return results - - -def _check_version() -> dict[str, Any]: - """Check 4: CLI version information. - - Returns: - Check result dict with the current CLI version. - """ - return { - "check": "version", - "name": "CLI version", - "status": "pass", - "message": f"kbagent v{__version__}", - } - - -def _format_doctor_human(console: Console, data: dict[str, Any]) -> None: - """Render doctor check results as a Rich panel with colored status indicators.""" - checks = data.get("checks", []) - - lines = [] - for check in checks: - status = check["status"] - if status == "pass": - icon = "[bold green]PASS[/bold green]" - elif status == "fail": - icon = "[bold red]FAIL[/bold red]" - elif status == "warn": - icon = "[bold yellow]WARN[/bold yellow]" - else: - icon = "[dim]SKIP[/dim]" - - lines.append(f" {icon} {check['name']}: {check['message']}") - - summary = data.get("summary", {}) - total = summary.get("total", 0) - passed = summary.get("passed", 0) - failed = summary.get("failed", 0) - warnings = summary.get("warnings", 0) - - lines.append("") - summary_parts = [f"{total} checks"] - if passed: - summary_parts.append(f"[green]{passed} passed[/green]") - if failed: - summary_parts.append(f"[red]{failed} failed[/red]") - if warnings: - summary_parts.append(f"[yellow]{warnings} warnings[/yellow]") - lines.append(f" Summary: {', '.join(summary_parts)}") - - panel = Panel("\n".join(lines), title="kbagent doctor", expand=False) - console.print(panel) - - def doctor_command(ctx: typer.Context) -> None: """Run health checks on CLI configuration and project connectivity.""" formatter = get_formatter(ctx) - config_store = get_service(ctx, "config_store") - - # Determine client factory - use the default unless we're in a test context - client_factory: ClientFactory = ctx.obj.get("client_factory", default_client_factory) - - all_checks: list[dict[str, Any]] = [] - - # Check 1: Config file exists with correct permissions - file_check = _check_config_file(config_store) - all_checks.append(file_check) - - # Check 2: Config file is valid JSON and parseable - valid_check, config = _check_config_valid(config_store) - all_checks.append(valid_check) - - # Check 3: Project connectivity - connectivity_checks = _check_connectivity(config, client_factory) - all_checks.extend(connectivity_checks) - - # Check 4: CLI version - version_check = _check_version() - all_checks.append(version_check) - - # Check 5: MCP server availability - mcp_service: McpService = ctx.obj.get("mcp_service", McpService(config_store)) - mcp_check = mcp_service.check_server_available() - all_checks.append(mcp_check) - - # Build summary - total = len(all_checks) - passed = sum(1 for c in all_checks if c["status"] == "pass") - failed = sum(1 for c in all_checks if c["status"] == "fail") - warnings = sum(1 for c in all_checks if c["status"] == "warn") - skipped = sum(1 for c in all_checks if c["status"] == "skip") - - result = { - "checks": all_checks, - "summary": { - "total": total, - "passed": passed, - "failed": failed, - "warnings": warnings, - "skipped": skipped, - "healthy": failed == 0, - }, - } - - formatter.output(result, _format_doctor_human) + doctor_service = get_service(ctx, "doctor_service") + result = doctor_service.run_checks() + formatter.output(result, format_doctor_panel) diff --git a/src/keboola_agent_cli/output.py b/src/keboola_agent_cli/output.py index 70d3062c..1a63977e 100644 --- a/src/keboola_agent_cli/output.py +++ b/src/keboola_agent_cli/output.py @@ -666,3 +666,39 @@ def _render_linked_buckets_table(console: Console, linked_buckets: list[dict[str console.print(table) console.print() + + +def format_doctor_panel(console: Console, data: dict[str, Any]) -> None: + """Render doctor check results as a Rich panel with colored status indicators. + + Args: + console: Rich Console instance. + data: Dict with "checks" list and "summary" dict from DoctorService. + """ + status_icons = { + "pass": "[bold green]PASS[/bold green]", + "fail": "[bold red]FAIL[/bold red]", + "warn": "[bold yellow]WARN[/bold yellow]", + } + + checks = data.get("checks", []) + lines = [ + f" {status_icons.get(c['status'], '[dim]SKIP[/dim]')} {c['name']}: {c['message']}" + for c in checks + ] + + summary = data.get("summary", {}) + parts = [f"{summary.get('total', 0)} checks"] + if summary.get("passed"): + parts.append(f"[green]{summary['passed']} passed[/green]") + if summary.get("failed"): + parts.append(f"[red]{summary['failed']} failed[/red]") + if summary.get("warnings"): + parts.append(f"[yellow]{summary['warnings']} warnings[/yellow]") + + lines.append("") + lines.append(f" Summary: {', '.join(parts)}") + + from rich.panel import Panel + + console.print(Panel("\n".join(lines), title="kbagent doctor", expand=False)) diff --git a/src/keboola_agent_cli/services/doctor_service.py b/src/keboola_agent_cli/services/doctor_service.py new file mode 100644 index 00000000..68a7ed8c --- /dev/null +++ b/src/keboola_agent_cli/services/doctor_service.py @@ -0,0 +1,253 @@ +"""Doctor service - health check logic for CLI configuration and connectivity. + +Runs checks for: +1. Config file existence and permissions (0600) +2. Config file valid JSON and parseable +3. Token verification for each project (API call with response time) +4. CLI version +5. MCP server availability + +Extracted from commands/doctor.py to respect the 3-layer architecture. +""" + +import json +import os +import stat +import time +from typing import Any + +from .. import __version__ +from ..config_store import ConfigStore +from ..errors import KeboolaApiError +from ..models import AppConfig +from .base import ClientFactory, default_client_factory +from .mcp_service import McpService + + +class DoctorService: + """Business logic for health checks. + + Accepts ConfigStore, client_factory, and McpService via DI + for easy testing with mocks. + """ + + def __init__( + self, + config_store: ConfigStore, + client_factory: ClientFactory | None = None, + mcp_service: McpService | None = None, + ) -> None: + self._config_store = config_store + self._client_factory = client_factory or default_client_factory + self._mcp_service = mcp_service or McpService(config_store) + + def run_checks(self) -> dict[str, Any]: + """Run all health checks and return structured results. + + Returns: + Dict with 'checks' list and 'summary' dict. + """ + all_checks: list[dict[str, Any]] = [] + + # Check 1: Config file exists with correct permissions + file_check = self._check_config_file() + all_checks.append(file_check) + + # Check 2: Config file is valid JSON and parseable + valid_check, config = self._check_config_valid() + all_checks.append(valid_check) + + # Check 3: Project connectivity + connectivity_checks = self._check_connectivity(config) + all_checks.extend(connectivity_checks) + + # Check 4: CLI version + version_check = self._check_version() + all_checks.append(version_check) + + # Check 5: MCP server availability + mcp_check = self._mcp_service.check_server_available() + all_checks.append(mcp_check) + + # Build summary + total = len(all_checks) + passed = sum(1 for c in all_checks if c["status"] == "pass") + failed = sum(1 for c in all_checks if c["status"] == "fail") + warnings = sum(1 for c in all_checks if c["status"] == "warn") + skipped = sum(1 for c in all_checks if c["status"] == "skip") + + return { + "checks": all_checks, + "summary": { + "total": total, + "passed": passed, + "failed": failed, + "warnings": warnings, + "skipped": skipped, + "healthy": failed == 0, + }, + } + + def _check_config_file(self) -> dict[str, Any]: + """Check 1: Config file exists and has correct permissions (0600). + + Returns: + Dict with check name, status (pass/fail/warn), and message. + """ + config_path = self._config_store.config_path + + if not config_path.exists(): + return { + "check": "config_file", + "name": "Config file", + "status": "warn", + "message": f"Config file not found at {config_path}. Run 'kbagent project add' to create it.", + } + + # Check permissions (Unix only) + try: + file_stat = os.stat(config_path) + mode = stat.S_IMODE(file_stat.st_mode) + if mode != 0o600: + return { + "check": "config_file", + "name": "Config file", + "status": "warn", + "message": f"Config file exists at {config_path} but has permissions {oct(mode)} (expected 0o600).", + } + except OSError: + # On platforms where permission checking is not reliable + pass + + return { + "check": "config_file", + "name": "Config file", + "status": "pass", + "message": f"Config file exists at {config_path} with correct permissions.", + } + + def _check_config_valid(self) -> tuple[dict[str, Any], AppConfig | None]: + """Check 2: Config file is valid JSON and parseable. + + Returns: + Tuple of (check result dict, parsed AppConfig or None on failure). + """ + config_path = self._config_store.config_path + + if not config_path.exists(): + return { + "check": "config_valid", + "name": "Config parseable", + "status": "skip", + "message": "No config file to validate.", + }, None + + try: + raw = config_path.read_text(encoding="utf-8") + except OSError as exc: + return { + "check": "config_valid", + "name": "Config parseable", + "status": "fail", + "message": f"Cannot read config file: {exc}", + }, None + + try: + json.loads(raw) + except json.JSONDecodeError as exc: + return { + "check": "config_valid", + "name": "Config parseable", + "status": "fail", + "message": f"Config file is not valid JSON: {exc}", + }, None + + try: + config = self._config_store.load() + except Exception as exc: + return { + "check": "config_valid", + "name": "Config parseable", + "status": "fail", + "message": f"Config file has invalid structure: {exc}", + }, None + + project_count = len(config.projects) + return { + "check": "config_valid", + "name": "Config parseable", + "status": "pass", + "message": f"Config file is valid JSON with {project_count} project(s).", + }, config + + def _check_connectivity( + self, + config: AppConfig | None, + ) -> list[dict[str, Any]]: + """Check 3: For each project, verify token via API call. + + Returns: + List of check result dicts, one per project. + """ + if config is None or not config.projects: + return [ + { + "check": "connectivity", + "name": "Project connectivity", + "status": "skip", + "message": "No projects configured.", + } + ] + + results = [] + for alias, project in config.projects.items(): + client = self._client_factory(project.stack_url, project.token) + start_time = time.monotonic() + try: + token_info = client.verify_token() + elapsed = time.monotonic() - start_time + results.append( + { + "check": "connectivity", + "name": f"Project '{alias}'", + "status": "pass", + "message": ( + f"Connected to {project.stack_url} " + f"(project: {token_info.project_name}, id: {token_info.project_id}) " + f"in {round(elapsed * 1000)}ms" + ), + "alias": alias, + "response_time_ms": round(elapsed * 1000), + } + ) + except KeboolaApiError as exc: + elapsed = time.monotonic() - start_time + results.append( + { + "check": "connectivity", + "name": f"Project '{alias}'", + "status": "fail", + "message": f"Failed: {exc.message}", + "alias": alias, + "error_code": exc.error_code, + "response_time_ms": round(elapsed * 1000), + } + ) + finally: + client.close() + + return results + + @staticmethod + def _check_version() -> dict[str, Any]: + """Check 4: CLI version information. + + Returns: + Check result dict with the current CLI version. + """ + return { + "check": "version", + "name": "CLI version", + "status": "pass", + "message": f"kbagent v{__version__}", + } diff --git a/src/keboola_agent_cli/services/lineage_service.py b/src/keboola_agent_cli/services/lineage_service.py index de9a87eb..46ef3667 100644 --- a/src/keboola_agent_cli/services/lineage_service.py +++ b/src/keboola_agent_cli/services/lineage_service.py @@ -2,10 +2,9 @@ Queries the Storage API for bucket sharing metadata and builds a graph of data flow edges between projects. Fetches buckets from -all projects in parallel using ThreadPoolExecutor. +all projects in parallel using BaseService._run_parallel(). """ -from concurrent.futures import ThreadPoolExecutor, as_completed from typing import Any from ..errors import KeboolaApiError @@ -20,7 +19,7 @@ class LineageService(BaseService): classifies buckets as shared or linked, and builds a deduplicated set of data flow edges. - Uses dependency injection for config_store and client_factory. + Uses BaseService._run_parallel() for concurrent project fetching. """ def _fetch_project_buckets( @@ -61,8 +60,8 @@ def get_lineage(self, aliases: list[str] | None = None) -> dict[str, Any]: """Analyze cross-project data lineage via bucket sharing. For each resolved project, fetches buckets with linkedBuckets info - in parallel using ThreadPoolExecutor, classifies them as shared or - linked, and builds deduplicated edges. + in parallel using BaseService._run_parallel(), classifies them as + shared or linked, and builds deduplicated edges. Args: aliases: Project aliases to query. None means all projects. @@ -80,7 +79,7 @@ def get_lineage(self, aliases: list[str] | None = None) -> dict[str, Any]: """ projects = self.resolve_projects(aliases) - # Early return for zero projects (avoids ThreadPoolExecutor(max_workers=0)) + # Early return for zero projects if not projects: return { "edges": [], @@ -104,46 +103,21 @@ def get_lineage(self, aliases: list[str] | None = None) -> dict[str, Any]: all_shared_buckets: list[dict[str, Any]] = [] all_linked_buckets: list[dict[str, Any]] = [] edges_by_key: dict[tuple[int, str, int, str], dict[str, Any]] = {} - errors: list[dict[str, str]] = [] - - # Fetch buckets from all projects in parallel - max_workers = min(len(projects), self._resolve_max_workers()) - with ThreadPoolExecutor(max_workers=max_workers) as executor: - future_to_alias = { - executor.submit(self._fetch_project_buckets, alias, project): alias - for alias, project in projects.items() - } - for future in as_completed(future_to_alias): - try: - result = future.result() - except Exception as exc: - # Unexpected exception from the future itself - proj_alias = future_to_alias[future] - errors.append( - { - "project_alias": proj_alias, - "error_code": "UNEXPECTED_ERROR", - "message": str(exc), - } - ) - continue - - # Distinguish success (3-tuple) from error (2-tuple) - if len(result) == 3: - alias, project, buckets = result - self._process_buckets( - buckets=buckets, - alias=alias, - project=project, - project_id_to_alias=project_id_to_alias, - all_shared_buckets=all_shared_buckets, - all_linked_buckets=all_linked_buckets, - edges_by_key=edges_by_key, - ) - else: - _alias, error_dict = result - errors.append(error_dict) + # Fetch buckets from all projects in parallel using BaseService._run_parallel() + successes, errors = self._run_parallel(projects, self._fetch_project_buckets) + + for success in successes: + alias, project, buckets = success + self._process_buckets( + buckets=buckets, + alias=alias, + project=project, + project_id_to_alias=project_id_to_alias, + all_shared_buckets=all_shared_buckets, + all_linked_buckets=all_linked_buckets, + edges_by_key=edges_by_key, + ) # Sort results for deterministic output (parallel execution order varies). # Use str() in sort keys: source_project_id comes from Pydantic (int) but diff --git a/src/keboola_agent_cli/services/mcp_service.py b/src/keboola_agent_cli/services/mcp_service.py index 046bd52a..0832b35d 100644 --- a/src/keboola_agent_cli/services/mcp_service.py +++ b/src/keboola_agent_cli/services/mcp_service.py @@ -535,6 +535,42 @@ def _call_read_tool( return asyncio.run(self._gather_read_results(projects, tool_name, tool_input)) + @staticmethod + async def _gather_results( + tasks: dict[str, "asyncio.Task[dict[str, Any]]"], + ) -> dict[str, Any]: + """Gather results from async tasks using asyncio.gather. + + Shared helper for both read and auto-expand gather operations. + Uses asyncio.gather with return_exceptions=True for true concurrency. + + Args: + tasks: Dict mapping project alias to asyncio.Task. + + Returns: + Dict with "results" list and "errors" list. + """ + aliases = list(tasks.keys()) + outcomes = await asyncio.gather(*tasks.values(), return_exceptions=True) + + all_results: list[dict[str, Any]] = [] + errors: list[dict[str, str]] = [] + + for alias, outcome in zip(aliases, outcomes): + if isinstance(outcome, BaseException): + errors.append( + { + "project_alias": alias, + "error_code": "MCP_ERROR", + "message": str(outcome), + } + ) + else: + outcome["project_alias"] = alias + all_results.append(outcome) + + return {"results": all_results, "errors": errors} + async def _gather_auto_expand_results( self, projects: dict[str, ProjectConfig], @@ -552,25 +588,7 @@ async def _gather_auto_expand_results( tasks[a] = asyncio.create_task( _connect_and_auto_expand(project, tool_name, tool_input, expand_config) ) - - all_results: list[dict[str, Any]] = [] - errors: list[dict[str, str]] = [] - - for a, task in tasks.items(): - try: - result = await task - result["project_alias"] = a - all_results.append(result) - except Exception as exc: - errors.append( - { - "project_alias": a, - "error_code": "MCP_ERROR", - "message": str(exc), - } - ) - - return {"results": all_results, "errors": errors} + return await self._gather_results(tasks) async def _gather_read_results( self, @@ -584,25 +602,7 @@ async def _gather_read_results( tasks[a] = asyncio.create_task( _connect_and_call_tool(project, tool_name, tool_input) ) - - all_results: list[dict[str, Any]] = [] - errors: list[dict[str, str]] = [] - - for a, task in tasks.items(): - try: - result = await task - result["project_alias"] = a - all_results.append(result) - except Exception as exc: - errors.append( - { - "project_alias": a, - "error_code": "MCP_ERROR", - "message": str(exc), - } - ) - - return {"results": all_results, "errors": errors} + return await self._gather_results(tasks) def check_server_available(self) -> dict[str, Any]: """Check if MCP server is available (for doctor command). diff --git a/src/keboola_agent_cli/services/org_service.py b/src/keboola_agent_cli/services/org_service.py index 849ef270..24c925e3 100644 --- a/src/keboola_agent_cli/services/org_service.py +++ b/src/keboola_agent_cli/services/org_service.py @@ -153,7 +153,14 @@ def setup_organization( "token": mask_token(registered.token) if registered else "***", "action": "added", }) - except (KeboolaApiError, Exception) as exc: + except KeboolaApiError as exc: + failed.append({ + "project_id": project_id, + "project_name": project_name, + "alias": alias, + "error": str(exc), + }) + except Exception as exc: failed.append({ "project_id": project_id, "project_name": project_name, diff --git a/tests/test_cli.py b/tests/test_cli.py index 1cb3ae38..8c039306 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -2127,7 +2127,7 @@ def test_doctor_connectivity_with_mock_client(self, tmp_path: Path) -> None: with ( patch("keboola_agent_cli.cli.ConfigStore") as MockStore, - patch("keboola_agent_cli.commands.doctor.default_client_factory") as MockFactory, + patch("keboola_agent_cli.services.doctor_service.default_client_factory") as MockFactory, ): MockStore.return_value = store MockFactory.return_value = mock_client @@ -2169,7 +2169,7 @@ def test_doctor_connectivity_failure(self, tmp_path: Path) -> None: with ( patch("keboola_agent_cli.cli.ConfigStore") as MockStore, - patch("keboola_agent_cli.commands.doctor.default_client_factory") as MockFactory, + patch("keboola_agent_cli.services.doctor_service.default_client_factory") as MockFactory, ): MockStore.return_value = store MockFactory.return_value = fail_client diff --git a/tests/test_doctor_service.py b/tests/test_doctor_service.py new file mode 100644 index 00000000..bd32e1b0 --- /dev/null +++ b/tests/test_doctor_service.py @@ -0,0 +1,397 @@ +"""Tests for DoctorService - health check logic extracted from doctor command.""" + +import json +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + +from keboola_agent_cli.config_store import ConfigStore +from keboola_agent_cli.errors import KeboolaApiError +from keboola_agent_cli.models import ProjectConfig, TokenVerifyResponse +from keboola_agent_cli.services.doctor_service import DoctorService +from keboola_agent_cli.services.mcp_service import McpService + + +def _make_mock_client( + project_name: str = "Test Project", + project_id: int = 1234, +) -> MagicMock: + """Create a mock KeboolaClient with verify_token returning valid data.""" + mock_client = MagicMock() + mock_client.verify_token.return_value = TokenVerifyResponse( + token_id="12345", + token_description="My Token", + project_id=project_id, + project_name=project_name, + owner_name=project_name, + ) + return mock_client + + +def _make_failing_client(error: KeboolaApiError) -> MagicMock: + """Create a mock KeboolaClient whose verify_token raises the given error.""" + mock_client = MagicMock() + mock_client.verify_token.side_effect = error + return mock_client + + +def _make_mcp_service_mock(status: str = "pass") -> MagicMock: + """Create a mock McpService returning a check_server_available result.""" + mock_mcp = MagicMock(spec=McpService) + mock_mcp.check_server_available.return_value = { + "check": "mcp_server", + "name": "MCP server", + "status": status, + "message": "MCP server available" if status == "pass" else "MCP server not found", + } + return mock_mcp + + +class TestDoctorServiceCheckConfigFile: + """Tests for DoctorService._check_config_file() - config file existence and permissions.""" + + def test_config_file_not_found(self, tmp_config_dir: Path) -> None: + """When config file does not exist, returns 'warn' status.""" + store = ConfigStore(config_dir=tmp_config_dir) + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + result = service._check_config_file() + + assert result["check"] == "config_file" + assert result["status"] == "warn" + assert "not found" in result["message"] + + def test_config_file_exists_correct_permissions(self, tmp_config_dir: Path) -> None: + """When config file exists with 0600 permissions, returns 'pass' status.""" + store = ConfigStore(config_dir=tmp_config_dir) + store.add_project( + "test", + ProjectConfig( + stack_url="https://connection.keboola.com", + token="901-xxx-testtoken1234", + project_name="Test", + project_id=1234, + ), + ) + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + result = service._check_config_file() + + assert result["check"] == "config_file" + assert result["status"] == "pass" + assert "correct permissions" in result["message"] + + def test_config_file_wrong_permissions(self, tmp_config_dir: Path) -> None: + """When config file exists with wrong permissions, returns 'warn' status.""" + store = ConfigStore(config_dir=tmp_config_dir) + store.add_project( + "test", + ProjectConfig( + stack_url="https://connection.keboola.com", + token="901-xxx-testtoken1234", + project_name="Test", + project_id=1234, + ), + ) + # Change permissions to 0644 + store.config_path.chmod(0o644) + + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + result = service._check_config_file() + + assert result["check"] == "config_file" + assert result["status"] == "warn" + assert "permissions" in result["message"] + assert "0o644" in result["message"] + + +class TestDoctorServiceCheckConfigValid: + """Tests for DoctorService._check_config_valid() - config file validation.""" + + def test_no_config_file_returns_skip(self, tmp_config_dir: Path) -> None: + """When config file does not exist, returns 'skip' status.""" + store = ConfigStore(config_dir=tmp_config_dir) + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + result, config = service._check_config_valid() + + assert result["check"] == "config_valid" + assert result["status"] == "skip" + assert config is None + + def test_valid_config_file(self, tmp_config_dir: Path) -> None: + """When config file is valid JSON with projects, returns 'pass'.""" + store = ConfigStore(config_dir=tmp_config_dir) + store.add_project( + "test", + ProjectConfig( + stack_url="https://connection.keboola.com", + token="901-xxx-testtoken1234", + project_name="Test", + project_id=1234, + ), + ) + + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + result, config = service._check_config_valid() + + assert result["check"] == "config_valid" + assert result["status"] == "pass" + assert "1 project" in result["message"] + assert config is not None + assert len(config.projects) == 1 + + def test_invalid_json_config(self, tmp_config_dir: Path) -> None: + """When config file contains invalid JSON, returns 'fail'.""" + store = ConfigStore(config_dir=tmp_config_dir) + config_path = tmp_config_dir / "config.json" + config_path.write_text("not valid json {{{", encoding="utf-8") + config_path.chmod(0o600) + + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + result, config = service._check_config_valid() + + assert result["check"] == "config_valid" + assert result["status"] == "fail" + assert "not valid JSON" in result["message"] + assert config is None + + def test_valid_json_invalid_structure(self, tmp_config_dir: Path) -> None: + """When config file is valid JSON but has invalid structure, returns 'fail'.""" + store = ConfigStore(config_dir=tmp_config_dir) + config_path = tmp_config_dir / "config.json" + # Valid JSON, but invalid structure for AppConfig + config_path.write_text('{"projects": "not-a-dict"}', encoding="utf-8") + config_path.chmod(0o600) + + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + result, config = service._check_config_valid() + + assert result["check"] == "config_valid" + assert result["status"] == "fail" + assert "invalid structure" in result["message"] + assert config is None + + +class TestDoctorServiceCheckConnectivity: + """Tests for DoctorService._check_connectivity() - API connectivity checks.""" + + def test_no_config_returns_skip(self, tmp_config_dir: Path) -> None: + """When config is None, returns skip for connectivity.""" + store = ConfigStore(config_dir=tmp_config_dir) + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + results = service._check_connectivity(None) + + assert len(results) == 1 + assert results[0]["check"] == "connectivity" + assert results[0]["status"] == "skip" + + def test_successful_connectivity(self, tmp_config_dir: Path) -> None: + """When API responds successfully, returns 'pass' with response time.""" + store = ConfigStore(config_dir=tmp_config_dir) + store.add_project( + "prod", + ProjectConfig( + stack_url="https://connection.keboola.com", + token="901-xxx-testtoken1234", + project_name="Production", + project_id=1234, + ), + ) + config = store.load() + + mock_client = _make_mock_client(project_name="Production", project_id=1234) + service = DoctorService( + config_store=store, + client_factory=lambda url, token: mock_client, + mcp_service=_make_mcp_service_mock(), + ) + + results = service._check_connectivity(config) + + assert len(results) == 1 + assert results[0]["check"] == "connectivity" + assert results[0]["status"] == "pass" + assert "Production" in results[0]["message"] + assert "response_time_ms" in results[0] + mock_client.close.assert_called_once() + + def test_connectivity_failure(self, tmp_config_dir: Path) -> None: + """When API call fails, returns 'fail' with error details.""" + store = ConfigStore(config_dir=tmp_config_dir) + store.add_project( + "bad", + ProjectConfig( + stack_url="https://connection.keboola.com", + token="901-badtoken-abcdef1234", + project_name="Bad", + project_id=9999, + ), + ) + config = store.load() + + error = KeboolaApiError( + message="Invalid token", + status_code=401, + error_code="INVALID_TOKEN", + retryable=False, + ) + fail_client = _make_failing_client(error) + service = DoctorService( + config_store=store, + client_factory=lambda url, token: fail_client, + mcp_service=_make_mcp_service_mock(), + ) + + results = service._check_connectivity(config) + + assert len(results) == 1 + assert results[0]["check"] == "connectivity" + assert results[0]["status"] == "fail" + assert "Invalid token" in results[0]["message"] + assert results[0]["error_code"] == "INVALID_TOKEN" + fail_client.close.assert_called_once() + + def test_multiple_projects_mixed(self, tmp_config_dir: Path) -> None: + """With multiple projects, each gets its own connectivity check.""" + store = ConfigStore(config_dir=tmp_config_dir) + store.add_project( + "prod", + ProjectConfig( + stack_url="https://connection.keboola.com", + token="901-xxx-testtoken1234", + project_name="Production", + project_id=1234, + ), + ) + store.add_project( + "bad", + ProjectConfig( + stack_url="https://connection.keboola.com", + token="901-badtoken-abcdef1234", + project_name="Bad", + project_id=9999, + ), + ) + config = store.load() + + error = KeboolaApiError( + message="Forbidden", + status_code=403, + error_code="ACCESS_DENIED", + retryable=False, + ) + + def factory(url: str, token: str) -> MagicMock: + if "badtoken" in token: + return _make_failing_client(error) + return _make_mock_client(project_name="Production", project_id=1234) + + service = DoctorService( + config_store=store, + client_factory=factory, + mcp_service=_make_mcp_service_mock(), + ) + + results = service._check_connectivity(config) + + assert len(results) == 2 + statuses = {r["alias"]: r["status"] for r in results} + assert statuses["prod"] == "pass" + assert statuses["bad"] == "fail" + + +class TestDoctorServiceCheckVersion: + """Tests for DoctorService._check_version() - CLI version check.""" + + def test_version_check_passes(self, tmp_config_dir: Path) -> None: + """Version check always returns 'pass' with version string.""" + store = ConfigStore(config_dir=tmp_config_dir) + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + result = service._check_version() + + assert result["check"] == "version" + assert result["status"] == "pass" + assert "kbagent v" in result["message"] + + +class TestDoctorServiceRunChecks: + """Tests for DoctorService.run_checks() - full health check orchestration.""" + + def test_run_checks_returns_all_checks_and_summary(self, tmp_config_dir: Path) -> None: + """run_checks returns a complete structure with checks and summary.""" + store = ConfigStore(config_dir=tmp_config_dir) + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + result = service.run_checks() + + assert "checks" in result + assert "summary" in result + assert len(result["checks"]) >= 4 # file, valid, connectivity, version, mcp + assert "total" in result["summary"] + assert "passed" in result["summary"] + assert "failed" in result["summary"] + assert "warnings" in result["summary"] + assert "skipped" in result["summary"] + assert "healthy" in result["summary"] + + def test_run_checks_no_config_is_healthy(self, tmp_config_dir: Path) -> None: + """With no config file, run_checks is still healthy (no failures).""" + store = ConfigStore(config_dir=tmp_config_dir) + service = DoctorService(config_store=store, mcp_service=_make_mcp_service_mock()) + + result = service.run_checks() + + assert result["summary"]["healthy"] is True + assert result["summary"]["failed"] == 0 + + def test_run_checks_with_connectivity_failure_is_unhealthy(self, tmp_config_dir: Path) -> None: + """When a connectivity check fails, healthy is False.""" + store = ConfigStore(config_dir=tmp_config_dir) + store.add_project( + "bad", + ProjectConfig( + stack_url="https://connection.keboola.com", + token="901-badtoken-abcdef1234", + project_name="Bad", + project_id=9999, + ), + ) + + error = KeboolaApiError( + message="Invalid token", + status_code=401, + error_code="INVALID_TOKEN", + retryable=False, + ) + fail_client = _make_failing_client(error) + service = DoctorService( + config_store=store, + client_factory=lambda url, token: fail_client, + mcp_service=_make_mcp_service_mock(), + ) + + result = service.run_checks() + + assert result["summary"]["healthy"] is False + assert result["summary"]["failed"] >= 1 + + def test_run_checks_includes_mcp_check(self, tmp_config_dir: Path) -> None: + """run_checks includes the MCP server availability check.""" + store = ConfigStore(config_dir=tmp_config_dir) + mock_mcp = _make_mcp_service_mock(status="warn") + service = DoctorService(config_store=store, mcp_service=mock_mcp) + + result = service.run_checks() + + mcp_checks = [c for c in result["checks"] if c["check"] == "mcp_server"] + assert len(mcp_checks) == 1 + assert mcp_checks[0]["status"] == "warn" + mock_mcp.check_server_available.assert_called_once()