diff --git a/strix/runtime/docker_connection.py b/strix/runtime/docker_connection.py index aa821aa9..3560faab 100644 --- a/strix/runtime/docker_connection.py +++ b/strix/runtime/docker_connection.py @@ -33,7 +33,7 @@ class DockerEndpoint: @property def label(self) -> str: - return f"{self.host or 'default socket'} ({self.source})" + return f"{self.host} ({self.source})" if self.host else self.source class DockerConnectionError(RuntimeError): @@ -59,22 +59,26 @@ def resolve_docker_endpoint(environ: dict[str, str] | None = None) -> DockerEndp return DockerEndpoint(host, "DOCKER_HOST") name = get_current_context_name() - if name != DEFAULT_CONTEXT: - try: - context = ContextAPI.get_context(name) - except Exception: # noqa: BLE001 - a broken context file must not hide Docker itself - context = None - if context is not None and context.Host: - return DockerEndpoint(context.Host, f"docker context '{name}'", context.TLSConfig) - return DockerEndpoint(None, "default socket") + if name == DEFAULT_CONTEXT: + return DockerEndpoint(None, "default socket") + + source = f"docker context '{name}'" + try: + context = ContextAPI.get_context(name) + except DockerException as exc: + raise DockerConnectionError(DockerEndpoint(None, source), exc) from exc + if context is None: + missing = DockerException(f"{source} is selected but does not exist") + raise DockerConnectionError(DockerEndpoint(None, source), missing) + return DockerEndpoint(context.Host, source, context.TLSConfig) def connect_docker() -> Any: """Return a ``docker.DockerClient`` for the resolved endpoint or raise DockerConnectionError.""" endpoint = resolve_docker_endpoint() try: - if endpoint.host is None: - return docker.from_env() + if endpoint.source == "DOCKER_HOST" or endpoint.host is None: + return docker.from_env() # also reads DOCKER_TLS_VERIFY and DOCKER_CERT_PATH return docker.DockerClient(base_url=endpoint.host, tls=endpoint.tls or False) except DockerException as exc: raise DockerConnectionError(endpoint, exc) from exc diff --git a/tests/test_docker_connection.py b/tests/test_docker_connection.py index 7253b2fb..5ea73dc8 100644 --- a/tests/test_docker_connection.py +++ b/tests/test_docker_connection.py @@ -51,17 +51,64 @@ def test_current_context_is_used_like_the_cli(monkeypatch: pytest.MonkeyPatch) - assert endpoint.source == "docker context 'desktop-linux'" -def test_default_and_broken_contexts_fall_back_to_the_sdk(monkeypatch: pytest.MonkeyPatch) -> None: +def test_default_context_uses_the_sdk_default(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(docker_connection, "get_current_context_name", lambda: "default") assert resolve_docker_endpoint({}) == DockerEndpoint(None, "default socket") - monkeypatch.setattr(docker_connection, "get_current_context_name", lambda: "gone") - def boom(_name: str) -> SimpleNamespace: - raise ValueError("bad meta.json") +def test_missing_or_broken_current_context_is_an_error(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(docker_connection, "get_current_context_name", lambda: "gone") + monkeypatch.setattr( + "strix.runtime.docker_connection.ContextAPI.get_context", lambda _name: None + ) + with pytest.raises( + DockerConnectionError, match="docker context 'gone' is selected but does not exist" + ): + resolve_docker_endpoint({}) + + def boom(_name: str) -> None: + raise DockerException("bad meta.json") monkeypatch.setattr("strix.runtime.docker_connection.ContextAPI.get_context", boom) - assert resolve_docker_endpoint({}) == DockerEndpoint(None, "default socket") + with pytest.raises(DockerConnectionError, match="bad meta") as info: + resolve_docker_endpoint({}) + assert info.value.endpoint.label == "docker context 'gone'" + + +def test_docker_host_goes_through_from_env_so_tls_settings_apply( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + docker_connection, + "resolve_docker_endpoint", + lambda _environ=None: DockerEndpoint("tcp://10.0.0.5:2376", "DOCKER_HOST"), + ) + client = object() + monkeypatch.setattr("strix.runtime.docker_connection.docker.from_env", lambda: client) + + def unexpected(**_kwargs: Any) -> None: + raise AssertionError("DOCKER_HOST must not bypass from_env") + + monkeypatch.setattr("strix.runtime.docker_connection.docker.DockerClient", unexpected) + assert docker_connection.connect_docker() is client + + +def test_context_endpoint_keeps_its_tls_config(monkeypatch: pytest.MonkeyPatch) -> None: + tls = object() + monkeypatch.setattr( + docker_connection, + "resolve_docker_endpoint", + lambda _environ=None: DockerEndpoint("tcp://remote:2376", "docker context 'remote'", tls), + ) + seen: dict[str, Any] = {} + + def client(**kwargs: Any) -> str: + seen.update(kwargs) + return "client" + + monkeypatch.setattr("strix.runtime.docker_connection.docker.DockerClient", client) + assert docker_connection.connect_docker() == "client" + assert seen == {"base_url": "tcp://remote:2376", "tls": tls} @pytest.mark.parametrize( @@ -78,7 +125,7 @@ def test_connect_docker_surfaces_the_root_cause( monkeypatch.setattr( docker_connection, "resolve_docker_endpoint", - lambda _environ=None: DockerEndpoint("unix:///nope.sock", "DOCKER_HOST"), + lambda _environ=None: DockerEndpoint("unix:///nope.sock", "docker context 'dead'"), ) def dead_client(**_kwargs: Any) -> None: @@ -117,7 +164,7 @@ def test_check_docker_connection_prints_endpoint_and_error_then_exits( assert exit_info.value.code == 1 assert reported == [("docker_unavailable", PermissionError)] assert "DOCKER NOT AVAILABLE" in out - assert "default socket (default socket)" in out + assert "Cannot connect to Docker at default socket." in out assert "PermissionError: [Errno 13] Permission denied" in out assert "docker info" in out