From 976782ca365f214df41a5e30e5ac7d0df401123e Mon Sep 17 00:00:00 2001 From: John Reeves Date: Wed, 16 Sep 2026 12:46:34 +1000 Subject: [PATCH] fix(auth): trustworthy s3 checks and non-blocking ssh probe (0.27.0) (#190) `dt auth check` misdiagnosed a real R2 credential problem three ways: 1. S3 endpoints used `aws sts get-caller-identity`, which Cloudflare R2 does not implement, so healthy endpoints reported "credentials check timed out". The probe also never used the per-repo AWS profile, so valid credentials in a named `~/.aws/credentials` profile read as "credentials not configured". 2. Every failure collapsed to "not configured", so a rejected secret (`SignatureDoesNotMatch`) sent the diagnosis looking for a missing profile instead of a wrong one. 3. A named SSH remote went through DVC's own (non-BatchMode, un-timed) filesystem, which blocks on a username prompt in a non-interactive shell / CI. Rewrite `_check_s3` to resolve the configured profile + endpoint URL and issue a bounded boto3 `list_objects_v2(MaxKeys=1)`, classifying the outcome as missing / rejected (code surfaced) / bucket-not-found / unreachable. Skip the DVC-native probe for SSH remotes when stdin is not a tty, falling back to the bounded `_check_ssh`. Co-Authored-By: Claude Opus 4.8 (1M context) --- dt/__init__.py | 2 +- dt/auth/__init__.py | 3 + dt/auth/checks.py | 206 ++++++++++++++++++++++++++++++++-------- pyproject.toml | 2 +- tests/unit/test_auth.py | 190 ++++++++++++++++++++++++++++++++---- 5 files changed, 345 insertions(+), 58 deletions(-) diff --git a/dt/__init__.py b/dt/__init__.py index 10ff1a5..6c2b12b 100644 --- a/dt/__init__.py +++ b/dt/__init__.py @@ -1,3 +1,3 @@ """DVC Tools - Convenient tools for working with DVC in HPC environments.""" -__version__ = "0.26.0" +__version__ = "0.27.0" diff --git a/dt/auth/__init__.py b/dt/auth/__init__.py index 2578575..2719031 100644 --- a/dt/auth/__init__.py +++ b/dt/auth/__init__.py @@ -132,6 +132,9 @@ _check_ssh, _check_ssh_remote_dir, _extract_ssh_remote_path, + _profile_from_source, + _resolve_s3_remote_settings, + _split_s3_url, _CHECKERS, _extract_remote_name, _get_owner_info, diff --git a/dt/auth/checks.py b/dt/auth/checks.py index 6b74bb5..aa631a5 100644 --- a/dt/auth/checks.py +++ b/dt/auth/checks.py @@ -4,6 +4,7 @@ import json import os import subprocess +import sys from dataclasses import dataclass, field from pathlib import Path from typing import Dict, List, Optional, Set, Tuple @@ -350,72 +351,190 @@ def _extract_remote_name(source: str) -> Optional[str]: return m.group(1) if m else None -def _check_s3(ep: Endpoint) -> CheckResult: - """Check an S3-compatible endpoint.""" - import shutil +# S3 error codes that mean credentials were *presented and rejected* (a wrong +# or stale secret), as opposed to being absent. Surfaced verbatim so a rejected +# secret no longer reads as "not configured" (issue #190). +_S3_REJECTED_CODES = { + 'SignatureDoesNotMatch', 'InvalidAccessKeyId', 'AccessDenied', + 'InvalidToken', 'ExpiredToken', 'TokenRefreshRequired', + 'AuthorizationHeaderMalformed', 'Forbidden', '403', +} - if not shutil.which('aws'): - return CheckResult( - endpoint=ep, status=STATUS_SKIP, - summary='aws CLI not installed', - hints=['Install the AWS CLI: pip install awscli'], - ) +# Codes that mean the credentials worked but the bucket/prefix is wrong. +_S3_MISSING_BUCKET_CODES = {'NoSuchBucket', 'NotFound', '404'} + + +def _profile_from_source(source: str) -> Optional[str]: + """AWS profile name implied by an endpoint's *source* string. + + Per-repo credentials use one profile per repo, named after the repo + (see :func:`dt.auth.credentials.configure_remotes`). An import child's + source reads ``DVC remote 'x' of ``; the owning repo is the + profile. Returns *None* for a top-level remote (no ``of``). + """ + import re + m = re.search(r' of ([^\s(]+)', source) + return m.group(1) if m else None + + +def _resolve_s3_remote_settings(ep: Endpoint) -> Tuple[Optional[str], Optional[str]]: + """Best-effort ``(endpoint_url, profile)`` for an S3 endpoint. + + Looks in the local ``.dvc/config`` first — by DVC remote name, then by + URL — then falls back to the profile-per-repo naming convention for + import children whose config lives in another repository. + """ + endpoint_url: Optional[str] = None + profile: Optional[str] = None - endpoint_url = None remote_name = _extract_remote_name(ep.source) - if remote_name: + if remote_name and ' of ' not in ep.source: endpoint_url = _get_dvc_remote_config(remote_name, 'endpointurl') + profile = _get_dvc_remote_config(remote_name, 'profile') + + if endpoint_url is None or profile is None: + try: + from .credentials import _get_project_s3_remotes + for values in _get_project_s3_remotes().values(): + if values.get('url') == ep.url: + endpoint_url = endpoint_url or values.get('endpointurl') + profile = profile or values.get('profile') + break + except Exception: + pass + + if profile is None: + profile = _profile_from_source(ep.source) + + return endpoint_url, profile + + +def _split_s3_url(url: str) -> Tuple[str, str]: + """Split ``s3://bucket/prefix`` into ``(bucket, prefix)``.""" + rest = url[len('s3://'):] if url.startswith('s3://') else url + bucket, _, prefix = rest.partition('/') + return bucket, prefix - extra_args: List[str] = [] - if endpoint_url: - extra_args = ['--endpoint-url', endpoint_url] +def _check_s3(ep: Endpoint) -> CheckResult: + """Check an S3-compatible endpoint with a bounded, profile-aware probe. + + Resolves the DVC remote's configured AWS profile and endpoint URL, then + issues a lightweight ``list_objects_v2(MaxKeys=1)`` with a per-call + timeout. This avoids ``aws sts get-caller-identity`` (which Cloudflare R2 + does not implement, so it timed out on healthy endpoints) and, unlike the + old default-credential CLI probe, actually uses the per-repo profile. The + outcome distinguishes *missing* credentials from *rejected* ones from an + *unreachable* endpoint (issue #190). + """ try: - cred_result = subprocess.run( - ['aws', 'sts', 'get-caller-identity'] + extra_args, - capture_output=True, text=True, timeout=10, + import boto3 + from botocore.config import Config as BotoConfig + from botocore.exceptions import ( + ClientError, ConnectTimeoutError, EndpointConnectionError, + NoCredentialsError, ProfileNotFound, ReadTimeoutError, ) - except (subprocess.TimeoutExpired, OSError): + except ImportError: return CheckResult( - endpoint=ep, status=STATUS_FAIL, - summary='credentials check timed out', + endpoint=ep, status=STATUS_SKIP, + summary='boto3 not installed', + hints=['Install boto3 to check S3 credentials: pip install boto3'], + ) + + endpoint_url, profile = _resolve_s3_remote_settings(ep) + profile_label = f"profile '{profile}'" if profile else 'default credentials' + + # Resolve credentials up front so a missing/unknown profile is reported as + # "missing", not as a downstream request error. + try: + session = ( + boto3.Session(profile_name=profile) if profile else boto3.Session() ) + creds = session.get_credentials() + except ProfileNotFound: + creds = None + session = None - if cred_result.returncode != 0: - hint = 'Configure AWS credentials in ~/.aws/credentials or environment variables' - if endpoint_url: - hint += f' for endpoint {endpoint_url}' + if creds is None: return CheckResult( endpoint=ep, status=STATUS_FAIL, - summary='credentials not configured', - hints=[hint], + summary=f'credentials missing — no {profile_label}', + hints=[ + f'No AWS {profile_label} found. Install it: ' + f'dt auth credentials install', + ], ) - bucket_prefix = ep.url + boto_config = BotoConfig( + connect_timeout=5, read_timeout=10, + retries={'max_attempts': 1, 'mode': 'standard'}, + ) + client = session.client('s3', endpoint_url=endpoint_url, config=boto_config) + + bucket, prefix = _split_s3_url(ep.url) try: - ls_result = subprocess.run( - ['aws', 's3', 'ls', bucket_prefix] + extra_args, - capture_output=True, text=True, timeout=15, + client.list_objects_v2(Bucket=bucket, Prefix=prefix, MaxKeys=1) + except (ConnectTimeoutError, ReadTimeoutError, EndpointConnectionError): + target = endpoint_url or 'AWS S3' + return CheckResult( + endpoint=ep, status=STATUS_FAIL, + summary='endpoint unreachable / timed out', + hints=[f'Check network access to the S3 endpoint: {target}'], ) - except (subprocess.TimeoutExpired, OSError): + except NoCredentialsError: return CheckResult( - endpoint=ep, status=STATUS_WARN, - summary='credentials OK, bucket check timed out', + endpoint=ep, status=STATUS_FAIL, + summary=f'credentials missing — no {profile_label}', + hints=['Install credentials: dt auth credentials install'], ) - - if ls_result.returncode != 0: + except ClientError as exc: + code = str(exc.response.get('Error', {}).get('Code', '')) or 'Unknown' + if code in _S3_REJECTED_CODES: + return CheckResult( + endpoint=ep, status=STATUS_FAIL, + summary=f'credentials rejected ({code})', + hints=[ + f'The {profile_label} was presented but rejected ({code}). ' + f'The secret is likely wrong or stale — reinstall it: ' + f'dt auth credentials install', + ], + ) + if code in _S3_MISSING_BUCKET_CODES: + return CheckResult( + endpoint=ep, status=STATUS_FAIL, + summary=f'bucket not found ({code})', + hints=[f'Check the bucket name/URL: {ep.url}'], + ) return CheckResult( endpoint=ep, status=STATUS_FAIL, - summary='credentials OK, bucket not accessible', - hints=[f'Check bucket exists and your credentials have access: {bucket_prefix}'], + summary=f'access failed ({code})', + hints=[_s3_manual_probe_hint(ep.url, profile, endpoint_url)], + ) + except Exception as exc: # defensive: unexpected boto/network failure + return CheckResult( + endpoint=ep, status=STATUS_FAIL, + summary=f'credentials check failed: {exc}', + hints=[_s3_manual_probe_hint(ep.url, profile, endpoint_url)], ) return CheckResult( endpoint=ep, status=STATUS_PASS, - summary='credentials OK, bucket accessible', + summary=f'credentials OK ({profile_label}), bucket accessible', ) +def _s3_manual_probe_hint( + url: str, profile: Optional[str], endpoint_url: Optional[str], +) -> str: + """A copy-pasteable ``aws s3 ls`` command for manual reproduction.""" + cmd = f'aws s3 ls {url}' + if profile: + cmd += f' --profile {profile}' + if endpoint_url: + cmd += f' --endpoint-url {endpoint_url}' + return f'Test manually: {cmd}' + + def _check_gs(ep: Endpoint) -> CheckResult: """Check a GCS endpoint.""" import shutil @@ -974,7 +1093,16 @@ def _try_check(ep: Endpoint, verbose: bool = False, remote_name = _extract_remote_name(ep.source) - if remote_name and ' of ' not in ep.source: + # The DVC-native probe opens the remote through DVC's own filesystem. For + # SSH that filesystem is not BatchMode and has no timeout, so in a + # non-interactive shell it blocks indefinitely on a username/password + # prompt (issue #190). Fall back to the bounded _check_ssh (BatchMode + + # ConnectTimeout) whenever stdin is not a tty. + use_dvc_native = bool(remote_name) and ' of ' not in ep.source + if ep.type == 'ssh' and not sys.stdin.isatty(): + use_dvc_native = False + + if use_dvc_native: dvc_result = _check_dvc_remote(ep, remote_name, verbose=verbose) if dvc_result is not None: return dvc_result diff --git a/pyproject.toml b/pyproject.toml index 7e24fb4..74cff3f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "dvc-tools" -version = "0.26.0" +version = "0.27.0" description = "Convenient tools for working with DVC in HPC environments" readme = "README.md" requires-python = ">=3.8" diff --git a/tests/unit/test_auth.py b/tests/unit/test_auth.py index 2f4ec22..ce72b69 100644 --- a/tests/unit/test_auth.py +++ b/tests/unit/test_auth.py @@ -76,6 +76,9 @@ _check_gs, _check_dvc_remote, _check_dvc_remote_impl, + _profile_from_source, + _resolve_s3_remote_settings, + _split_s3_url, _try_check, _extract_remote_name, _short_repo_name, @@ -1086,32 +1089,160 @@ def test_curl_not_installed(self, _): # _check_s3 tests # ============================================================================= +class TestSplitS3Url: + """Tests for _split_s3_url.""" + + def test_bucket_and_prefix(self): + assert _split_s3_url('s3://bucket/a/b') == ('bucket', 'a/b') + + def test_bucket_only(self): + assert _split_s3_url('s3://bucket') == ('bucket', '') + + def test_no_scheme(self): + assert _split_s3_url('bucket/prefix') == ('bucket', 'prefix') + + +class TestProfileFromSource: + """Tests for _profile_from_source.""" + + def test_import_child(self): + assert _profile_from_source("DVC remote 'gadi' of chromium (default)") == 'chromium' + + def test_import_child_no_default(self): + assert _profile_from_source("DVC remote 'x' of visium") == 'visium' + + def test_top_level_none(self): + assert _profile_from_source("DVC remote 'cloud' (default)") is None + + +class TestResolveS3RemoteSettings: + """Tests for _resolve_s3_remote_settings.""" + + def test_top_level_reads_dvc_config(self): + ep = Endpoint(type='s3', url='s3://bucket', source="DVC remote 'cloud'") + with patch('dt.auth.checks._get_dvc_remote_config') as mock_cfg: + mock_cfg.side_effect = lambda name, key: { + 'endpointurl': 'https://r2.example.com', + 'profile': 'myrepo', + }.get(key) + endpoint_url, profile = _resolve_s3_remote_settings(ep) + assert endpoint_url == 'https://r2.example.com' + assert profile == 'myrepo' + + def test_import_child_uses_repo_name_convention(self): + ep = Endpoint(type='s3', url='s3://bucket', + source="DVC remote 'gadi' of chromium (default)") + # No local .dvc/config match, so the profile comes from the source repo. + with patch('dt.auth.credentials._get_project_s3_remotes', return_value={}): + endpoint_url, profile = _resolve_s3_remote_settings(ep) + assert profile == 'chromium' + + class TestCheckS3: - """Tests for _check_s3.""" + """Tests for _check_s3 (boto3-based, profile-aware — issue #190).""" + + def _session(self, *, creds=object(), list_side_effect=None, + list_return=None): + """Build a mock boto3 Session with a stubbed s3 client.""" + session = MagicMock() + session.get_credentials.return_value = creds + client = MagicMock() + if list_side_effect is not None: + client.list_objects_v2.side_effect = list_side_effect + else: + client.list_objects_v2.return_value = list_return or {'KeyCount': 0} + session.client.return_value = client + return session, client + + def test_missing_credentials(self): + ep = Endpoint(type='s3', url='s3://bucket', source='cloud') + session, _ = self._session(creds=None) + with patch('dt.auth.checks._resolve_s3_remote_settings', + return_value=(None, 'wts')), \ + patch('boto3.Session', return_value=session): + r = _check_s3(ep) + assert r.status == STATUS_FAIL + assert 'missing' in r.summary + assert "profile 'wts'" in r.summary - @patch('shutil.which', return_value=None) - def test_aws_not_installed(self, _): + def test_profile_not_found(self): + from botocore.exceptions import ProfileNotFound ep = Endpoint(type='s3', url='s3://bucket', source='cloud') - r = _check_s3(ep) - assert r.status == STATUS_SKIP + with patch('dt.auth.checks._resolve_s3_remote_settings', + return_value=(None, 'wts')), \ + patch('boto3.Session', + side_effect=ProfileNotFound(profile='wts')): + r = _check_s3(ep) + assert r.status == STATUS_FAIL + assert 'missing' in r.summary - @patch('shutil.which', return_value='/usr/bin/aws') - @patch('subprocess.run') - def test_credentials_fail(self, mock_run, _): - mock_run.return_value = MagicMock(returncode=1) + def test_rejected_signature_does_not_match(self): + from botocore.exceptions import ClientError ep = Endpoint(type='s3', url='s3://bucket', source='cloud') - r = _check_s3(ep) + err = ClientError( + {'Error': {'Code': 'SignatureDoesNotMatch'}}, 'ListObjectsV2') + session, _ = self._session(list_side_effect=err) + with patch('dt.auth.checks._resolve_s3_remote_settings', + return_value=('https://r2', 'wts')), \ + patch('boto3.Session', return_value=session): + r = _check_s3(ep) assert r.status == STATUS_FAIL - assert 'credentials not configured' in r.summary + assert 'rejected' in r.summary + assert 'SignatureDoesNotMatch' in r.summary - @patch('shutil.which', return_value='/usr/bin/aws') - @patch('subprocess.run') - def test_full_pass(self, mock_run, _): - mock_run.return_value = MagicMock(returncode=0) + def test_access_denied_is_rejected(self): + from botocore.exceptions import ClientError + ep = Endpoint(type='s3', url='s3://bucket', source='cloud') + err = ClientError({'Error': {'Code': 'AccessDenied'}}, 'ListObjectsV2') + session, _ = self._session(list_side_effect=err) + with patch('dt.auth.checks._resolve_s3_remote_settings', + return_value=('https://r2', 'wts')), \ + patch('boto3.Session', return_value=session): + r = _check_s3(ep) + assert r.status == STATUS_FAIL + assert 'rejected' in r.summary + + def test_endpoint_unreachable(self): + from botocore.exceptions import EndpointConnectionError + ep = Endpoint(type='s3', url='s3://bucket', source='cloud') + err = EndpointConnectionError(endpoint_url='https://r2') + session, _ = self._session(list_side_effect=err) + with patch('dt.auth.checks._resolve_s3_remote_settings', + return_value=('https://r2', 'wts')), \ + patch('boto3.Session', return_value=session): + r = _check_s3(ep) + assert r.status == STATUS_FAIL + assert 'unreachable' in r.summary or 'timed out' in r.summary + + def test_bucket_not_found(self): + from botocore.exceptions import ClientError + ep = Endpoint(type='s3', url='s3://bucket', source='cloud') + err = ClientError({'Error': {'Code': 'NoSuchBucket'}}, 'ListObjectsV2') + session, _ = self._session(list_side_effect=err) + with patch('dt.auth.checks._resolve_s3_remote_settings', + return_value=('https://r2', 'wts')), \ + patch('boto3.Session', return_value=session): + r = _check_s3(ep) + assert r.status == STATUS_FAIL + assert 'bucket not found' in r.summary + + def test_pass_uses_profile_and_endpoint(self): ep = Endpoint(type='s3', url='s3://bucket/prefix', source='cloud') - r = _check_s3(ep) + session, client = self._session() + with patch('dt.auth.checks._resolve_s3_remote_settings', + return_value=('https://r2', 'xenium')), \ + patch('boto3.Session', return_value=session) as mock_sess: + r = _check_s3(ep) assert r.status == STATUS_PASS - assert 'bucket accessible' in r.summary + assert "profile 'xenium'" in r.summary + # The resolved profile and endpoint URL are actually used. + mock_sess.assert_called_once_with(profile_name='xenium') + _, kwargs = session.client.call_args + assert kwargs['endpoint_url'] == 'https://r2' + _, list_kwargs = client.list_objects_v2.call_args + assert list_kwargs['Bucket'] == 'bucket' + assert list_kwargs['Prefix'] == 'prefix' + assert list_kwargs['MaxKeys'] == 1 # ============================================================================= @@ -1340,6 +1471,31 @@ def test_import_source_child_skips_dvc_native(self): result = _try_check(ep) mock_dvc.assert_not_called() + def test_ssh_skips_dvc_native_when_not_a_tty(self): + """A named SSH remote must not use the unbounded DVC-native probe + in a non-interactive shell — it would hang on a username prompt + (issue #190). It falls back to the bounded _check_ssh instead.""" + ep = Endpoint(type='ssh', url='ssh://host/path', + source="DVC remote 'nci'") + with patch('dt.auth.checks._check_dvc_remote') as mock_dvc, \ + patch('dt.auth.checks.sys.stdin.isatty', return_value=False), \ + patch('subprocess.run', return_value=MagicMock(returncode=0)): + result = _try_check(ep) + mock_dvc.assert_not_called() + assert result.endpoint is ep + + def test_ssh_uses_dvc_native_when_interactive(self): + """When stdin is a tty the DVC-native probe is still used for a + named SSH remote (interactive behaviour is preserved).""" + ep = Endpoint(type='ssh', url='ssh://host/path', + source="DVC remote 'nci'") + dvc_result = CheckResult(endpoint=ep, status=STATUS_PASS, summary='via DVC') + with patch('dt.auth.checks._check_dvc_remote', return_value=dvc_result) as mock_dvc, \ + patch('dt.auth.checks.sys.stdin.isatty', return_value=True): + result = _try_check(ep) + mock_dvc.assert_called_once() + assert result.summary == 'via DVC' + class TestCheckEndpoints: """Tests for check_endpoints."""